Skip to content

checkpoint

Module diagram

classDiagram
  class checkpoint {
  }
  class aes_crypto {
  }
  class api {
  }
  class checkpoint_handler {
  }
  class metadata {
  }
  class api {
  }
  class metadata {
  }
  checkpoint_handler --> aes_crypto
  checkpoint_handler --> metadata
  metadata --> metadata

checkpoint

체크포인트 파일을 다루기 위한 모듈입니다. Checkpoint는 버전별로 다른 구조를 가집니다.

CHECKPOINT_SPEC은 다음과 같은 구조를 가집니다.

CHECKPOINT_SPEC (Dict[ByteString, Dict]): {
    version (ByteString): {
        "version": ByteString,  # YYMMDDV
        "metadata_size": int,  # 메타데이터의 크기 (bytes)
        "weight_position": int,  # 체크포인트 파일에서 weight가 시작하는 위치 (bytes)
    },
    ...
}

CHECKPOINT_VERSION_LENGTH module-attribute

CHECKPOINT_VERSION_LENGTH = 7

CHECKPOINT_LAST_VERSION module-attribute

CHECKPOINT_LAST_VERSION = b'2509090'

CHECKPOINT_SPEC module-attribute

CHECKPOINT_SPEC = {b'2212210': {'version': b'2212210', 'metadata_size': 1024 * 1024, 'weight_position': CHECKPOINT_VERSION_LENGTH + 1024 * 1024}, b'2401040': {'version': b'2401040', 'metadata_size': 12 * 1024 * 1024, 'weight_position': CHECKPOINT_VERSION_LENGTH + 12 * 1024 * 1024}, b'2509090': {'version': b'2509090', 'metadata_size': 30 * 1024 * 1024, 'weight_position': CHECKPOINT_VERSION_LENGTH + 30 * 1024 * 1024}}

WEIGHT_POSITION module-attribute

WEIGHT_POSITION = None

aes_crypto

체크포인트 암호화 관련 함수들을 정의합니다.

torch.save로 저장 가능한 파이썬 오브젝트를 암호화하거나 복호화합니다.

  • encrypt: torch.save로 저장 가능한 파이썬 오브젝트를 암호화합니다.
  • decrypt: torch.load로 로드 가능한 바이트를 복호화합니다.
  • encrypt_dict: dict 자체 혹은 요소들을 암호화합니다.
  • decrypt_dict: dict를 복호화합니다.
  • encrypt_save_dict: dict를 암호화 후 torch.save 로 저장합니다.
  • decrypt_load_dict: torch.load로 파일 로드 후 복호화해 dict로 로드합니다.

BUFFER_SIZE module-attribute

BUFFER_SIZE = 64 * 1024

encrypt

encrypt(obj: Any, password: str, buffer_size: int = BUFFER_SIZE, **torch_save_args) -> bytes

torch.save로 picklable한 obj를 Aes 암호화 후 바이트로 리턴합니다.

Parameters:

  • obj (Any) –

    description

  • password (str) –

    description

  • buffer_size (int, default: BUFFER_SIZE ) –

    description. Defaults to BUFFER_SIZE.

Returns:

  • bytes ( bytes ) –

    description

Source code in SaigeToolkit/checkpoint/aes_crypto.py
def encrypt(obj: Any, password: str, buffer_size: int = BUFFER_SIZE, **torch_save_args) -> bytes:
    """`torch.save`로 picklable한 obj를 Aes 암호화 후 바이트로 리턴합니다.

    Args:
        obj (Any): _description_
        password (str): _description_
        buffer_size (int, optional): _description_. Defaults to BUFFER_SIZE.

    Returns:
        bytes: _description_
    """
    buffer = io.BytesIO()
    torch.save(obj, buffer, **torch_save_args)
    buffer.seek(0)
    buffer_encrypt = io.BytesIO()
    pyAesCrypt.encryptStream(buffer, buffer_encrypt, password, buffer_size)
    return buffer_encrypt.getvalue()

decrypt

decrypt(encrpyted: bytes, password: str, buffer_size: int = BUFFER_SIZE, **torch_load_args) -> Any

torch.load로 로드 가능한 바이트를 Aes 복호화 후 파이썬 오브젝트로 리턴합니다.

Parameters:

  • encrpyted (bytes) –

    description

  • password (str) –

    description

  • buffer_size (int, default: BUFFER_SIZE ) –

    description. Defaults to BUFFER_SIZE.

Returns:

  • Any ( Any ) –

    description

Source code in SaigeToolkit/checkpoint/aes_crypto.py
def decrypt(encrpyted: bytes, password: str, buffer_size: int = BUFFER_SIZE, **torch_load_args) -> Any:
    """`torch.load`로 로드 가능한 바이트를 Aes 복호화 후 파이썬 오브젝트로 리턴합니다.

    Args:
        encrpyted (bytes): _description_
        password (str): _description_
        buffer_size (int, optional): _description_. Defaults to BUFFER_SIZE.

    Returns:
        Any: _description_
    """
    decrypt_buffer = io.BytesIO()
    pyAesCrypt.decryptStream(
        io.BytesIO(encrpyted), decrypt_buffer, password, buffer_size, len(encrpyted)
    )
    decrypt_buffer.seek(0)
    return torch.load(decrypt_buffer, weights_only=False, **torch_load_args)

encrypt_dict

encrypt_dict(data: Mapping, password: Optional[str] = None, keys_encrypt: Optional[List[str]] = None, **torch_save_args)

dict 자체 혹은 요소들을 암호화합니다.

Parameters:

  • data (Mapping) –

    description

  • password (Optional[str], default: None ) –

    description. Defaults to None.

  • keys_encrypt (Optional[List[str]], default: None ) –

    None이 아닐 경우 keys_encrypt에 있는 키의 값들만 암호화합니다. Defaults to None.

Source code in SaigeToolkit/checkpoint/aes_crypto.py
def encrypt_dict(
    data: Mapping,
    password: Optional[str] = None,
    keys_encrypt: Optional[List[str]] = None,
    **torch_save_args,
):
    """dict 자체 혹은 요소들을 암호화합니다.

    Args:
        data (Mapping): _description_
        password (Optional[str], optional): _description_. Defaults to None.
        keys_encrypt (Optional[List[str]], optional): None이 아닐 경우 keys_encrypt에 있는 키의 값들만 암호화합니다. Defaults to None.
    """
    data_ = data.copy()
    if password is not None:
        if keys_encrypt is not None:
            for key in keys_encrypt:
                if key not in data_:
                    continue
                data_[key] = encrypt(data_[key], password, **torch_save_args)
        else:
            data_ = encrypt(data_, password, **torch_save_args)
    return data_

decrypt_dict

decrypt_dict(data, password: Optional[str] = None, keys_encrypt: Optional[List[str]] = None, **torch_load_args) -> Mapping

dict를 복호화합니다.

Parameters:

  • data (Any) –

    description

  • password (Optional[str], default: None ) –

    description. Defaults to None.

  • keys_encrypt (Optional[List[str]], default: None ) –

    None이 아닐 경우 keys_encrypt에 있는 키의 값들만 복호화합니다. Defaults to None.

Raises:

  • FileNotFoundError

    description

Returns:

  • Mapping ( Mapping ) –

    description

Source code in SaigeToolkit/checkpoint/aes_crypto.py
def decrypt_dict(
    data,
    password: Optional[str] = None,
    keys_encrypt: Optional[List[str]] = None,
    **torch_load_args,
) -> Mapping:
    """dict를 복호화합니다.

    Args:
        data (Any): _description_
        password (Optional[str], optional): _description_. Defaults to None.
        keys_encrypt (Optional[List[str]], optional): None이 아닐 경우 keys_encrypt에 있는 키의 값들만 복호화합니다. Defaults to None.

    Raises:
        FileNotFoundError: _description_

    Returns:
        Mapping: _description_
    """
    if password is not None:
        if keys_encrypt is not None:
            for key in keys_encrypt:
                if key not in data:
                    continue
                data[key] = decrypt(data[key], password, **torch_load_args)
        else:
            data = decrypt(data, password, **torch_load_args)
    return data

encrypt_save_dict

encrypt_save_dict(data: Mapping, path: str, password: Optional[str] = None, keys_encrypt: Optional[List[str]] = None, **torch_save_args) -> None

dict를 암호화 후 torch.save 로 저장합니다.

Parameters:

  • data (Mapping) –

    description

  • path (str) –

    description

  • password (Optional[str], default: None ) –

    description. Defaults to None.

  • keys_encrypt (Optional[List[str]], default: None ) –

    None이 아닐 경우 keys_encrypt에 있는 키의 값들만 암호화합니다. Defaults to None.

Source code in SaigeToolkit/checkpoint/aes_crypto.py
def encrypt_save_dict(
    data: Mapping,
    path: str,
    password: Optional[str] = None,
    keys_encrypt: Optional[List[str]] = None,
    **torch_save_args,
) -> None:
    """dict를 암호화 후 `torch.save` 로 저장합니다.

    Args:
        data (Mapping): _description_
        path (str): _description_
        password (Optional[str], optional): _description_. Defaults to None.
        keys_encrypt (Optional[List[str]], optional): None이 아닐 경우 keys_encrypt에 있는 키의 값들만 암호화합니다. Defaults to None.
    """
    data_ = encrypt_dict(data, password=password, keys_encrypt=keys_encrypt, **torch_save_args)
    os.makedirs(os.path.dirname(path), exist_ok=True)
    with open(path, "wb") as f:
        torch.save(data_, f, **torch_save_args)

decrypt_load_dict

decrypt_load_dict(path: str, password: Optional[str] = None, keys_encrypt: Optional[List[str]] = None, **torch_load_args) -> Mapping

torch.load로 파일 로드 후 복호화해 dict로 로드합니다.

Parameters:

  • path (str) –

    description

  • password (Optional[str], default: None ) –

    description. Defaults to None.

  • keys_encrypt (Optional[List[str]], default: None ) –

    None이 아닐 경우 keys_encrypt에 있는 키의 값들만 복호화합니다. Defaults to None.

Raises:

  • FileNotFoundError

    description

Returns:

  • Mapping ( Mapping ) –

    description

Source code in SaigeToolkit/checkpoint/aes_crypto.py
@handle_unpickling_error
def decrypt_load_dict(
    path: str,
    password: Optional[str] = None,
    keys_encrypt: Optional[List[str]] = None,
    **torch_load_args,
) -> Mapping:
    """`torch.load`로 파일 로드 후 복호화해 dict로 로드합니다.

    Args:
        path (str): _description_
        password (Optional[str], optional): _description_. Defaults to None.
        keys_encrypt (Optional[List[str]], optional): None이 아닐 경우 keys_encrypt에 있는 키의 값들만 복호화합니다. Defaults to None.

    Raises:
        FileNotFoundError: _description_

    Returns:
        Mapping: _description_
    """
    if not os.path.isfile(path):
        raise error.FileNotFoundError
    data = torch.load(path, weights_only=False, **torch_load_args)
    data = decrypt_dict(data, password=password, keys_encrypt=keys_encrypt, **torch_load_args)
    return data

api

error_handler module-attribute

error_handler = get_error_handler(SaigeToolkitError)

convert_checkpoint

convert_checkpoint(checkpoint_path: str, save_path: Optional[str] = None)

22년12월21일 이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

이 API는 추후 백엔드와의 논의를 통해 점진적으로 사용을 줄이고, 최종적으로는 삭제되어야 합니다. 현재는 torch와의 의존성을 줄이기 위해, 실제 호출 시에만 관련 패키지가 로드되도록 lazy import가 적용되어 있습니다.

Parameters:

  • checkpoint_path (str) –

    이전 버전 체크포인트 경로

  • save_path (Optional[str], default: None ) –

    저장할 경로. None인 경우 checkpoint_path를 덮어씁니다. Defaults to None.

Source code in SaigeToolkit/checkpoint/api.py
@error_handler
def convert_checkpoint(checkpoint_path: str, save_path: Optional[str] = None):
    """22년12월21일 이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

    이 API는 추후 백엔드와의 논의를 통해 점진적으로 사용을 줄이고, 최종적으로는 삭제되어야 합니다.
    현재는 torch와의 의존성을 줄이기 위해, 실제 호출 시에만 관련 패키지가 로드되도록 lazy import가 적용되어 있습니다.

    Args:
        checkpoint_path (str): 이전 버전 체크포인트 경로
        save_path (Optional[str], optional): 저장할 경로. None인 경우 `checkpoint_path`를 덮어씁니다. Defaults to None.
    """
    checkpoint_handler.convert_checkpoint(checkpoint_path=checkpoint_path, save_path=save_path)

checkpoint_handler

CheckpointHandler class를 정의합니다. CheckpointHandler는 체크포인트를 저장하고 불러오는 역할을 합니다.

CheckpointHandler는 다음과 같은 기능을 제공합니다. 1. 체크포인트 저장 2. 체크포인트 불러오기 3. 체크포인트 변환 4. 체크포인트 유효성 검사

CheckpointHandler는 다음과 같은 메서드를 제공합니다. 1. save: 체크포인트를 저장합니다. 2. load: 체크포인트를 불러옵니다. 3. convert: 이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 4. load_checkpoint: 체크포인트를 열어 로드합니다. 5. is_valid_checkpoint: 체크포인트가 유효한지 검증하고, 유효하면 True, 그렇지 않으면 False 리턴합니다.

logger module-attribute

logger = getLogger('SaigeResearch')

CheckpointHandler

_keys_encrypt class-attribute instance-attribute
_keys_encrypt = ['config', 'state_dict']
save classmethod
save(checkpoint_path: str, version: str, config: Dict, state_dict: Dict, metadata: Optional[Dict] = None, common: Optional[Dict] = None, password: Optional[str] = None) -> None
Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
@classmethod
def save(
    cls,
    checkpoint_path: str,
    version: str,
    config: Dict,
    state_dict: Dict,
    metadata: Optional[Dict] = None,
    common: Optional[Dict] = None,
    password: Optional[str] = None,
) -> None:
    metadata = metadata or {}
    common = common or {}
    checkpoint = {
        "version": version,
        "config": config,
        "common": common,
        "state_dict": state_dict,
    }
    checkpoint = encrypt_dict(
        data=checkpoint,
        password=password,
        keys_encrypt=cls._keys_encrypt,
    )

    os.makedirs(os.path.dirname(checkpoint_path), exist_ok=True)
    weight_position = CHECKPOINT_SPEC[CHECKPOINT_LAST_VERSION]["weight_position"]
    with open(checkpoint_path, "wb") as f:
        f.write(CHECKPOINT_LAST_VERSION)
        f.seek(weight_position)
        pickle.dump(checkpoint, f)

    write_metadata(checkpoint_path=checkpoint_path, metadata=metadata)
load classmethod
load(checkpoint_path: str, password: Optional[str] = None) -> Dict
Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
@classmethod
def load(
    cls,
    checkpoint_path: str,
    password: Optional[str] = None,
) -> Dict:
    if not os.path.isfile(checkpoint_path):
        raise ModelFileNotFoundError

    metadata = read_metadata(checkpoint_path=checkpoint_path)
    if metadata is False:
        logger.info(f"{type(cls).__name__} converting old checkpoint {checkpoint_path}")
        cls.convert(checkpoint_path=checkpoint_path)
        metadata = read_metadata(checkpoint_path=checkpoint_path)

    checkpoint = cls.load_checkpoint(checkpoint_path)

    checkpoint = decrypt_dict(
        data=checkpoint,
        password=password,
        keys_encrypt=cls._keys_encrypt,
        map_location="cpu",
    )
    checkpoint["metadata"] = metadata
    return checkpoint
convert classmethod
convert(checkpoint_path: str, save_path: Optional[str] = None)

이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

Parameters:

  • checkpoint_path (str) –

    이전 버전 체크포인트 경로

  • save_path (str, default: None ) –

    저장할 경로. None인 경우 checkpoint_path를 덮어씁니다. Defaults to None.

Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
@classmethod
def convert(cls, checkpoint_path: str, save_path: Optional[str] = None):
    """이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

    Args:
        checkpoint_path (str): 이전 버전 체크포인트 경로
        save_path (str): 저장할 경로. None인 경우 `checkpoint_path`를 덮어씁니다. Defaults to None.
    """
    if save_path is None:
        save_path = checkpoint_path
    checkpoint = decrypt_load_dict(
        path=checkpoint_path,
        password=None,
        map_location="cpu",
    )
    checkpoint.pop("metadata")
    cls.save(checkpoint_path=save_path, **checkpoint)
load_checkpoint classmethod
load_checkpoint(checkpoint_path: str) -> Dict

Open and load checkpoint from path

Parameters:

  • checkpoint_path (str) –

    checkpoint path

Raises:

  • InvalidModelFileError

    다음과 같은 경우 에러 레이즈

Returns:

  • Dict ( Dict ) –

    (일부 암호화된) checkpoint

Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
@classmethod
@handle_unpickling_error
def load_checkpoint(cls, checkpoint_path: str) -> Dict:
    """Open and load checkpoint from path

    Args:
        checkpoint_path (str): checkpoint path

    Raises:
        InvalidModelFileError: 다음과 같은 경우 에러 레이즈
        1) 잘못된 WEIGHT_POSITION이 사용된 경우
        2) 파일이 pickle파일이 아닌 경우
        3) 열린 checkpoint 내부의 형식이 올바르지 않은 경우

    Returns:
        Dict: (일부 암호화된) checkpoint
    """

    checkpoint_version = read_checkpoint_version(checkpoint_path)
    if checkpoint_version not in CHECKPOINT_SPEC:
        raise InvalidModelFileError
    weight_position = CHECKPOINT_SPEC[checkpoint_version]["weight_position"]

    with open(checkpoint_path, "rb") as f:
        f.seek(weight_position)
        checkpoint = pickle.load(f)

    if not cls.is_valid_checkpoint(checkpoint):
        raise InvalidModelFileError

    return checkpoint
is_valid_checkpoint classmethod
is_valid_checkpoint(checkpoint: Dict) -> bool

Checkpoint가 유효한지 검증하고, 유효하면 True, 그렇지 않으면 False 리턴

Note

현재 로직은 다음 세가지만을 체크함. 1) checkpoint가 python dictionary 인지 2) checkpoint가 모든 key를 가지고 있는지 3) "version" key에 해당하는 value가 유효한지

Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
@classmethod
def is_valid_checkpoint(cls, checkpoint: Dict) -> bool:
    """Checkpoint가 유효한지 검증하고, 유효하면 True, 그렇지 않으면 False 리턴

    Note:
        현재 로직은 다음 세가지만을 체크함.
        1) checkpoint가 python dictionary 인지
        2) checkpoint가 모든 key를 가지고 있는지
        3) "version" key에 해당하는 value가 유효한지
    """
    if not isinstance(checkpoint, dict):
        return False

    for k in ["version", "config", "common", "state_dict"]:
        if k not in checkpoint:
            return False

    try:
        _ = Version.from_string(checkpoint.get("version"))
    except VersionFormatError:
        return False

    return True

convert_checkpoint

convert_checkpoint(checkpoint_path: str, save_path: Optional[str] = None)

22년12월21일 이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

Parameters:

  • checkpoint_path (str) –

    이전 버전 체크포인트 경로

  • save_path (Optional[str], default: None ) –

    저장할 경로. None인 경우 checkpoint_path를 덮어씁니다. Defaults to None.

Source code in SaigeToolkit/checkpoint/checkpoint_handler.py
def convert_checkpoint(checkpoint_path: str, save_path: Optional[str] = None):
    """22년12월21일 이전 버전 체크포인트를 최신 구조로 변환해 저장합니다. 저장되어있던 metadata는 제거됩니다.

    Args:
        checkpoint_path (str): 이전 버전 체크포인트 경로
        save_path (Optional[str], optional): 저장할 경로. None인 경우 `checkpoint_path`를 덮어씁니다. Defaults to None.
    """
    CheckpointHandler.convert(checkpoint_path=checkpoint_path, save_path=save_path)

metadata

체크포인트 파일의 버전 및 메타데이터를 다루기 위한 모듈입니다.

api

체크포인트 파일의 버전 및 metadata를 읽고 쓰는 API를 제공합니다.

error_handler module-attribute
error_handler = get_error_handler(SaigeToolkitError)
read_checkpoint_version
read_checkpoint_version(checkpoint_path: str) -> ByteString

체크포인트 파일의 버전을 읽어옵니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

Returns:

  • ByteString ( ByteString ) –

    체크포인트 파일의 버전

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/api.py
@error_handler
def read_checkpoint_version(checkpoint_path: str) -> ByteString:
    """체크포인트 파일의 버전을 읽어옵니다.

    Args:
        checkpoint_path (str): checkpoint_path

    Returns:
        ByteString: 체크포인트 파일의 버전

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    return _read_checkpoint_version(checkpoint_path=checkpoint_path)
read_metadata
read_metadata(checkpoint_path: str) -> Union[Dict, bool]

체크포인트 파일의 metadata 섹션 데이터를 읽어옵니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

Returns:

  • Union[Dict, bool]

    Union[Dict, bool]: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/api.py
@error_handler
def read_metadata(checkpoint_path: str) -> Union[Dict, bool]:
    """체크포인트 파일의 metadata 섹션 데이터를 읽어옵니다.

    Args:
        checkpoint_path (str): checkpoint_path

    Returns:
        Union[Dict, bool]: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    return _read_metadata(checkpoint_path=checkpoint_path)
write_metadata
write_metadata(checkpoint_path: str, metadata: Dict) -> bool

체크포인트 파일의 metadata 섹션 데이터를 수정해 저장합니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

  • metadata (Dict) –

    metadata

Returns:

  • bool ( bool ) –

    체크포인트 버전이 맞지 않는 경우 False 값을 리턴

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/api.py
@error_handler
def write_metadata(checkpoint_path: str, metadata: Dict) -> bool:
    """체크포인트 파일의 metadata 섹션 데이터를 수정해 저장합니다.

    Args:
        checkpoint_path (str): checkpoint_path
        metadata (Dict): metadata

    Returns:
        bool: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    return _write_metadata(checkpoint_path=checkpoint_path, metadata=metadata)

metadata

체크포인트 파일의 버전 및 metadata를 읽고 쓰기 위한 메서드를 제공합니다.

read_checkpoint_version
read_checkpoint_version(checkpoint_path: str) -> ByteString

체크포인트 파일의 버전을 읽어옵니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

Returns:

  • ByteString ( ByteString ) –

    체크포인트 파일의 버전

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/metadata.py
@handle_unpickling_error
def read_checkpoint_version(checkpoint_path: str) -> ByteString:
    """체크포인트 파일의 버전을 읽어옵니다.

    Args:
        checkpoint_path (str): checkpoint_path

    Returns:
        ByteString: 체크포인트 파일의 버전

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    if not os.path.isfile(checkpoint_path):
        raise ModelFileNotFoundError

    with open(checkpoint_path, "rb") as f:
        return f.read(CHECKPOINT_VERSION_LENGTH)
read_metadata
read_metadata(checkpoint_path: str) -> Union[Dict, bool]

체크포인트 파일의 metadata 섹션 데이터를 읽어옵니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

Returns:

  • Union[Dict, bool]

    Union[Dict, bool]: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/metadata.py
@handle_unpickling_error
def read_metadata(checkpoint_path: str) -> Union[Dict, bool]:
    """체크포인트 파일의 metadata 섹션 데이터를 읽어옵니다.

    Args:
        checkpoint_path (str): checkpoint_path

    Returns:
        Union[Dict, bool]: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    if not os.path.isfile(checkpoint_path):
        raise ModelFileNotFoundError
    with open(checkpoint_path, "rb") as f:
        version = f.read(CHECKPOINT_VERSION_LENGTH)
        return False if version not in CHECKPOINT_SPEC else pickle.load(f)
write_metadata
write_metadata(checkpoint_path: str, metadata: Dict) -> bool

체크포인트 파일의 metadata 섹션 데이터를 수정해 저장합니다.

Parameters:

  • checkpoint_path (str) –

    checkpoint_path

  • metadata (Dict) –

    metadata

Returns:

  • bool ( bool ) –

    체크포인트 버전이 맞지 않는 경우 False 값을 리턴

Raises:

  • ModelFileNotFoundError

    checkpoint_path에 해당하는 파일이 없는 경우

Source code in SaigeToolkit/checkpoint/metadata/metadata.py
def write_metadata(checkpoint_path: str, metadata: Dict) -> bool:
    """체크포인트 파일의 metadata 섹션 데이터를 수정해 저장합니다.

    Args:
        checkpoint_path (str): checkpoint_path
        metadata (Dict): metadata

    Returns:
        bool: 체크포인트 버전이 맞지 않는 경우 False 값을 리턴

    Raises:
        ModelFileNotFoundError: checkpoint_path에 해당하는 파일이 없는 경우
    """
    if not os.path.isfile(checkpoint_path):
        raise ModelFileNotFoundError
    with open(checkpoint_path, "r+b") as f:
        version = f.read(CHECKPOINT_VERSION_LENGTH)
        if version not in CHECKPOINT_SPEC:
            return False
        bytes = pickle.dumps(metadata)
        metadata_size = CHECKPOINT_SPEC[version]["metadata_size"]
        if len(bytes) >= metadata_size:
            raise ValueError("metadata size exceeds the limit")
        f.write(bytes)
        return True