Skip to content

reproducibility

util.reproducibility

실험 재현성을 위한 유틸리티 함수를 제공합니다.

TIME_FORMAT module-attribute

TIME_FORMAT = '%Y-%m-%d %H:%M:%S KST'

set_randomness

set_randomness(seed=0, deterministic=False)
Source code in SaigeToolkit/util/reproducibility.py
def set_randomness(seed=0, deterministic=False):
    random.seed(seed)
    os.environ["PYTHONHASHSEED"] = str(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    torch.cuda.manual_seed(seed)
    torch.cuda.manual_seed_all(seed)
    if deterministic:
        torch.backends.cudnn.deterministic = True
        torch.backends.cudnn.benchmark = False

store_config

store_config(attr: Optional[str] = None)

클래스 메소드 호출 시 파라미터를 self.attr 변수에 dict로 저장하는 데코레이터를 리턴합니다.

Usage

class MyClass: @store_config(attr="_config") def init(self, param1, param2): pass

MyClass 인스턴스 생성 시 self._config에 파라미터가 저장됩니다.

Source code in SaigeToolkit/util/reproducibility.py
def store_config(attr: Optional[str] = None):
    """클래스 메소드 호출 시 파라미터를 `self.attr` 변수에 dict로 저장하는 데코레이터를 리턴합니다.

    Usage:
        class MyClass:
            @store_config(attr="_config")
            def __init__(self, param1, param2):
                pass

        MyClass 인스턴스 생성 시 self._config에 파라미터가 저장됩니다.

    """

    def _store_config(method, attr=attr):
        attr = attr or f"{method.__name__}_params"

        @wraps(method)
        def decorator(self, *args, **kwargs):
            signature = inspect.signature(method)
            bound_args = signature.bind(self, *args, **kwargs)
            bound_args.apply_defaults()
            config = dict(bound_args.arguments)
            config.pop("self")
            config = copy.deepcopy(config)
            setattr(self, attr, config)
            return method(self, *args, **kwargs)

        return decorator

    return _store_config

get_function_args

get_function_args() -> Dict

Returns the argument received by the parent function as a dictionary.

Returns:

  • Dict ( Dict ) –

    keyword arguments dictionary

Example

code: def func(A=1, B=2, C=3): print(get_function_args()) func(C=5)

results: {'A': 1, 'B': 2, 'C': 5}

Note

해당 함수를 호출하기 전에 argument를 수정하는 경우, 수정된 argument가 반환됩니다. Example: code: def func(A=1, B=2, C=3): B = 7 print(get_function_args()) func(C=5)

results:
    {'A': 1, 'B': 7, 'C': 5}
Source code in SaigeToolkit/util/reproducibility.py
def get_function_args() -> Dict:
    """Returns the argument received by the parent function as a dictionary.

    Returns:
        Dict: keyword arguments dictionary

    Example:
        code:
            def func(A=1, B=2, C=3):
                print(get_function_args())
            func(C=5)

        results:
            {'A': 1, 'B': 2, 'C': 5}

    Note:
        해당 함수를 호출하기 전에 argument를 수정하는 경우, 수정된 argument가 반환됩니다.
        Example:
            code:
                def func(A=1, B=2, C=3):
                    B = 7
                    print(get_function_args())
                func(C=5)

            results:
                {'A': 1, 'B': 7, 'C': 5}
    """
    frame = sys._getframe(1)
    args, _, _, values = inspect.getargvalues(frame)
    return copy.deepcopy({arg: values[arg] for arg in args})

get_source_git_info

get_source_git_info(git_path: Optional[str] = None) -> Dict

Parse git info

Parameters:

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

    git directory path. Defaults to None.

Returns:

  • dict ( Dict ) –

    { "directory": str, "commit": str, "modified files": List[str], "untracked files": List[str], "submodules_info": List[Dict],

  • Dict

    }

Source code in SaigeToolkit/util/reproducibility.py
def get_source_git_info(git_path: Optional[str] = None) -> Dict:
    """Parse git info

    Args:
        git_path (Optional[str], optional): git directory path. Defaults to None.

    Returns:
        dict: {
            "directory": str,
            "commit": str,
            "modified files": List[str],
            "untracked files": List[str],
            "submodules_info": List[Dict],
        }
    """
    if git is None:
        return {}

    def _get_source_git_info_by_repo(repo: git.Repo) -> Dict:
        info = {
            "directory": repo.working_dir,
            "commit": repo.head.commit.hexsha,
            "modified files": [item.a_path for item in repo.index.diff(None)],
            "untracked files": repo.untracked_files,
            "submodules_info": [
                _get_source_git_info_by_repo(submodule.module()) for submodule in repo.submodules
            ],
        }
        return info

    try:
        repo = git.Repo(git_path, search_parent_directories=git_path is None)
        return_value = _get_source_git_info_by_repo(repo)
    except git.GitError:
        return_value = {}

    return return_value

get_source_info

get_source_info(git_path: Optional[str] = None) -> Dict

Returns current time (KST) + git info + run command

Parameters:

  • git_path (str, default: None ) –

    git directory path. Defaults to GIT_PATH.

Returns:

  • dict ( Dict ) –

    {"time": str, "command": str, **git_info}

Source code in SaigeToolkit/util/reproducibility.py
def get_source_info(git_path: Optional[str] = None) -> Dict:
    """Returns current time (KST) + git info + run command

    Args:
        git_path (str, optional): git directory path. Defaults to GIT_PATH.

    Returns:
        dict: {"time": str, "command": str, **git_info}
    """
    return {
        "time": datetime.datetime.now(pytz.timezone("Asia/Seoul")).strftime(TIME_FORMAT),
        **get_source_git_info(git_path),
        "command": "python " + " ".join(sys.argv),
    }