Skip to content

feature_extractor

model.feature_extractor

_types module-attribute

_types = {__name__: _Ae8mfor _type in _types}

ForwardHookFeatureExtractor

ForwardHookFeatureExtractor(model: Module, layers: Dict[str, str])

Bases: Module

Extract intermediate features from the model with forward hooks.

return example: {"layer1": layer1_output, "layer2": layer2_output, ..}

Note
  • DOES NOT work with DataParallel. forward input이 단일 tensor가 아닌 경우 device로 분류하는 기존 방식이 동작하지 않기 때문에 제거함.
  • DOES work with DistributedDataParallel.
Source code in SaigeToolkit/model/feature_extractor.py
def __init__(
    self,
    model: nn.Module,
    layers: Dict[str, str],
):
    super().__init__()

    self.model = model
    self.layers = layers
    self._outputs = []

    for layer_name, output_name in self.layers.items():
        hook = self._get_hook_function(output_name)
        self.model.get_submodule(layer_name).register_forward_hook(hook)

model instance-attribute

model = model

layers instance-attribute

layers = layers

_outputs instance-attribute

_outputs = []

_get_hook_function

_get_hook_function(output_name: str)
Source code in SaigeToolkit/model/feature_extractor.py
def _get_hook_function(self, output_name: str):
    def _hook(module, input, output):
        if not self._outputs:
            self._outputs.append({})
        self._outputs[0][output_name] = output

    return _hook

forward

forward(input: Any) -> Dict[str, Any]
Source code in SaigeToolkit/model/feature_extractor.py
def forward(self, input: Any) -> Dict[str, Any]:
    _ = self.model(input)
    return self._outputs.pop()

build_backbone

build_backbone(_target_: str, **params) -> Module
Source code in SaigeToolkit/model/backbone/builder.py
def build_backbone(_target_: str, **params) -> nn.Module:
    if _target_ in native_models:
        builder = native_models[_target_]
    elif _target_.startswith("torchvision.models."):
        builder = getattr(torchvision.models, _target_.split(".")[-1])
    elif _target_ == "timm.create_model":
        import timm

        builder = timm.create_model
    elif _target_ == "torch.hub.load":
        builder = torch.hub.load
    else:
        raise ValueError(f"Unknown backbone target: {_target_}")

    return builder(**params)

build_feature_extractor

build_feature_extractor(backbone: Union[dict, Module], _target_: str = 'create_feature_extractor', layers: Optional[List[str]] = None, **params) -> Module

Build feature extractor

Parameters:

  • backbone (dict | Module) –

    backbone config.

  • _target_ (str, default: 'create_feature_extractor' ) –

    feature extractor target. Defaults to "create_feature_extractor".

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

    feature layers to extract. Defaults to None.

Returns:

  • Module

    nn.Module: feature extractor

Note
  • ForwardHookFeatureExtractor 의 경우 hook을 사용하므로 기존 모델과의 동일성이 보장되지만 불필요한 레이어까지 모두 사용하게됨.
  • IntermediateLayerGetter 의 경우 타겟 레이어 이후의 레이어를 제거해 효율적. 하지만 기존 모델의 forward가 복잡한 경우 동일성이 보장되지 않을 수 있음.
  • create_feature_extractor 의 경우 연산 그래프를 고려하여 타겟 레이어 이외의 레이어를 제거해 동일성과 효율성 측면에서 좋음. 하지만 timm 모델은 사용할 수 없는듯.
Source code in SaigeToolkit/model/feature_extractor.py
def build_feature_extractor(
    backbone: Union[dict, nn.Module],
    _target_: str = "create_feature_extractor",
    layers: Optional[List[str]] = None,
    **params,
) -> nn.Module:
    """Build feature extractor

    Args:
        backbone (dict | nn.Module): backbone config.
        _target_ (str, optional): feature extractor target. Defaults to "create_feature_extractor".
        layers (List[str], optional): feature layers to extract. Defaults to None.

    Returns:
        nn.Module: feature extractor

    Note:
        - `ForwardHookFeatureExtractor` 의 경우 hook을 사용하므로 기존 모델과의 동일성이 보장되지만 불필요한 레이어까지 모두 사용하게됨.
        - `IntermediateLayerGetter` 의 경우 타겟 레이어 이후의 레이어를 제거해 효율적. 하지만 기존 모델의 forward가 복잡한 경우 동일성이 보장되지 않을 수 있음.
        - `create_feature_extractor` 의 경우 연산 그래프를 고려하여 타겟 레이어 이외의 레이어를 제거해 동일성과 효율성 측면에서 좋음. 하지만 timm 모델은 사용할 수 없는듯.

    """
    if layers is None:
        layers = ["layer1", "layer2", "layer3"]

    if isinstance(backbone, dict):
        backbone = build_backbone(**backbone)
    return_nodes = {layer: layer for layer in layers}
    return _types[_target_](backbone, return_nodes, **params)