frbench 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
frbench/__init__.py ADDED
@@ -0,0 +1,58 @@
1
+ import os
2
+
3
+ os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")
4
+
5
+ from ._config import CACHE, RELEASE, REPO, __version__
6
+ from .fr import FR
7
+ from .types import FRDetectResult, FREmbedResult
8
+ from .utils.download import (
9
+ ModelInfo,
10
+ download_assets,
11
+ get_asset,
12
+ list_assets,
13
+ list_models,
14
+ refresh_manifest,
15
+ )
16
+ from .utils.log import (
17
+ FRBenchWarning,
18
+ download_verbose_enabled,
19
+ set_download_verbose,
20
+ set_verbose,
21
+ set_warnings,
22
+ warnings_enabled,
23
+ )
24
+ from .utils.update_check import set_update_check, update_check_enabled
25
+ from ._exceptions import (
26
+ FRBenchAssetNotFoundError,
27
+ FRBenchConfigError,
28
+ FRBenchDownloadError,
29
+ FRBenchError,
30
+ )
31
+
32
+ __all__ = [
33
+ "__version__",
34
+ "FR",
35
+ "FREmbedResult",
36
+ "FRDetectResult",
37
+ "ModelInfo",
38
+ "CACHE",
39
+ "REPO",
40
+ "RELEASE",
41
+ "FRBenchWarning",
42
+ "FRBenchError",
43
+ "FRBenchDownloadError",
44
+ "FRBenchAssetNotFoundError",
45
+ "FRBenchConfigError",
46
+ "set_warnings",
47
+ "set_download_verbose",
48
+ "set_verbose",
49
+ "warnings_enabled",
50
+ "download_verbose_enabled",
51
+ "set_update_check",
52
+ "update_check_enabled",
53
+ "get_asset",
54
+ "list_assets",
55
+ "list_models",
56
+ "download_assets",
57
+ "refresh_manifest",
58
+ ]
frbench/_config.py ADDED
@@ -0,0 +1,51 @@
1
+ """Runtime configuration for FRBench (env vars + module-level overrides)."""
2
+ from __future__ import annotations
3
+
4
+ import os
5
+ from typing import Optional
6
+
7
+ __version__ = "1.0.0"
8
+
9
+ _DEFAULT_CACHE = os.path.expanduser(os.path.expandvars("~/.frbench"))
10
+ _DEFAULT_REPO = "HKU-TASR/FRBench"
11
+ _DEFAULT_RELEASE = "weights-v1.0.0"
12
+
13
+ # Module-level overrides (set before first download, or via env vars).
14
+ CACHE: str = os.environ.get("FRBENCH_CACHE", _DEFAULT_CACHE)
15
+ REPO: str = os.environ.get("FRBENCH_REPO", _DEFAULT_REPO)
16
+ RELEASE: str = os.environ.get("FRBENCH_RELEASE", _DEFAULT_RELEASE)
17
+
18
+ _DETECTOR_PREFIX = "retinaface_"
19
+
20
+
21
+ def get_cache() -> str:
22
+ """Return the active cache directory."""
23
+ return os.path.expanduser(os.path.expandvars(os.environ.get("FRBENCH_CACHE", CACHE)))
24
+
25
+
26
+ def get_repo() -> str:
27
+ """Return the active GitHub repo slug."""
28
+ return os.environ.get("FRBENCH_REPO", REPO)
29
+
30
+
31
+ def get_release() -> str:
32
+ """Return the active GitHub release tag."""
33
+ return os.environ.get("FRBENCH_RELEASE", RELEASE)
34
+
35
+
36
+ def is_detector_asset(name: str) -> bool:
37
+ """True if *name* is a RetinaFace detector asset (not an FR model)."""
38
+ return name.startswith(_DETECTOR_PREFIX)
39
+
40
+
41
+ def parse_model_key(name: str) -> Optional[tuple[str, str, str]]:
42
+ """Parse a manifest model key into ``(backbone, loss, dataset)``.
43
+
44
+ Returns ``None`` for detector assets or malformed keys.
45
+ """
46
+ if is_detector_asset(name):
47
+ return None
48
+ parts = name.rsplit("_", 2)
49
+ if len(parts) != 3:
50
+ return None
51
+ return parts[0], parts[1], parts[2]
frbench/_exceptions.py ADDED
@@ -0,0 +1,17 @@
1
+ """FRBench exception hierarchy."""
2
+
3
+
4
+ class FRBenchError(Exception):
5
+ """Base class for FRBench errors."""
6
+
7
+
8
+ class FRBenchDownloadError(FRBenchError):
9
+ """Raised when an asset cannot be downloaded or verified."""
10
+
11
+
12
+ class FRBenchAssetNotFoundError(FRBenchDownloadError):
13
+ """Raised when a requested asset key is missing from the manifest."""
14
+
15
+
16
+ class FRBenchConfigError(FRBenchError):
17
+ """Raised when model or detector configuration is invalid."""
@@ -0,0 +1,154 @@
1
+ from typing import Dict, Any
2
+ import copy
3
+
4
+ import torch
5
+ from torch import nn
6
+
7
+ from .resnet import *
8
+ from .resnetv2 import *
9
+ from .ir import *
10
+ from .irse import *
11
+ from .densenet import *
12
+ from .efficientnet import *
13
+ from .mobilefacenet import *
14
+ from .mobilenet import *
15
+ from .convnext import *
16
+ from .swin_v1 import *
17
+ from .swin_v2 import *
18
+ from .swin_mlp import *
19
+ from .mobilevit_v1 import *
20
+ from .mobilevit_v2 import *
21
+ from .mobilevit_v3 import *
22
+
23
+ BACKBONE_REGISTRY = {
24
+ # ResNet Variants
25
+ 'resnet-18': ResNet_18,
26
+ 'resnet-34': ResNet_34,
27
+ 'resnet-50': ResNet_50,
28
+ 'resnet-100': ResNet_100,
29
+ 'resnet-101': ResNet_101,
30
+ 'resnet-152': ResNet_152,
31
+ 'resnet-200': ResNet_200,
32
+ # ResNetV2 Variants
33
+ 'resnetv2-18': ResNetV2_18,
34
+ 'resnetv2-34': ResNetV2_34,
35
+ 'resnetv2-50': ResNetV2_50,
36
+ 'resnetv2-100': ResNetV2_100,
37
+ 'resnetv2-101': ResNetV2_101,
38
+ 'resnetv2-152': ResNetV2_152,
39
+ 'resnetv2-200': ResNetV2_200,
40
+ # IR Variants
41
+ 'ir-18': IR_18,
42
+ 'ir-34': IR_34,
43
+ 'ir-50': IR_50,
44
+ 'ir-100': IR_100,
45
+ 'ir-101': IR_101,
46
+ 'ir-152': IR_152,
47
+ 'ir-200': IR_200,
48
+ # IR-SE Variants
49
+ 'irse-18': IR_SE_18,
50
+ 'irse-34': IR_SE_34,
51
+ 'irse-50': IR_SE_50,
52
+ 'irse-100': IR_SE_100,
53
+ 'irse-101': IR_SE_101,
54
+ 'irse-152': IR_SE_152,
55
+ 'irse-185': IR_SE_185,
56
+ 'irse-200': IR_SE_200,
57
+ # DenseNet Variants
58
+ 'densenet-121': DenseNet_121,
59
+ 'densenet-169': DenseNet_169,
60
+ 'densenet-201': DenseNet_201,
61
+ # EfficientNetV1 Variants
62
+ 'efficientnetv1-b0': EfficientNetV1_B0,
63
+ 'efficientnetv1-b1': EfficientNetV1_B1,
64
+ 'efficientnetv1-b2': EfficientNetV1_B2,
65
+ 'efficientnetv1-b3': EfficientNetV1_B3,
66
+ 'efficientnetv1-b4': EfficientNetV1_B4,
67
+ 'efficientnetv1-b5': EfficientNetV1_B5,
68
+ 'efficientnetv1-b6': EfficientNetV1_B6,
69
+ 'efficientnetv1-b7': EfficientNetV1_B7,
70
+ # MobileNet Variants
71
+ 'mobilenet-w1': MobileNet_W1,
72
+ 'mobilenet-w3d4': MobileNet_W3D4,
73
+ 'mobilenet-wd2': MobileNet_WD2,
74
+ 'mobilenet-wd4': MobileNet_WD4,
75
+ # MobileNetV2 Variants
76
+ 'mobilenetv2-w1': MobileNetV2_W1,
77
+ 'mobilenetv2-w3d4': MobileNetV2_W3D4,
78
+ 'mobilenetv2-wd2': MobileNetV2_WD2,
79
+ 'mobilenetv2-wd4': MobileNetV2_WD4,
80
+ # MobileNetV3 Variants
81
+ 'mobilenetv3-l': MobileNetV3_Large,
82
+ 'mobilenetv3-s': MobileNetV3_Small,
83
+ # MobileNetV4 Variants
84
+ 'mobilenetv4conv-s': MobileNetV4_Conv_Small,
85
+ 'mobilenetv4conv-m': MobileNetV4_Conv_Medium,
86
+ 'mobilenetv4conv-l': MobileNetV4_Conv_Large,
87
+ 'mobilenetv4hybrid-m': MobileNetV4_Hybrid_Medium,
88
+ 'mobilenetv4hybrid-l': MobileNetV4_Hybrid_Large,
89
+ # MobileFaceNet Variants
90
+ 'mobilefacenet': MobileFaceNet_ECA,
91
+ 'mobilefacenet-plain': MobileFaceNet_Plain,
92
+ 'convnext-t': ConvNeXt_Tiny,
93
+ 'convnext-s': ConvNeXt_Small,
94
+ 'convnext-b': ConvNeXt_Base,
95
+ 'convnext-l': ConvNeXt_Large,
96
+ 'convnext-xl': ConvNeXt_XLarge,
97
+ 'convnextv2-atto': ConvNeXtV2_Atto,
98
+ 'convnextv2-femto': ConvNeXtV2_Femto,
99
+ 'convnextv2-pico': ConvNeXtV2_Pico,
100
+ 'convnextv2-n': ConvNeXtV2_Nano,
101
+ 'convnextv2-t': ConvNeXtV2_Tiny,
102
+ 'convnextv2-s': ConvNeXtV2_Small,
103
+ 'convnextv2-b': ConvNeXtV2_Base,
104
+ 'convnextv2-l': ConvNeXtV2_Large,
105
+ 'convnextv2-xl': ConvNeXtV2_Huge,
106
+ # SwinV1 Variants
107
+ 'swinv1-t': SwinV1_Tiny, # 224 recommended; 112 works with 3 stages instead of 4
108
+ 'swinv1-s': SwinV1_Small,
109
+ 'swinv1-b': SwinV1_Base,
110
+ 'swinv1-l': SwinV1_Large,
111
+ # SwinMLP Variants (experimental spatial MLP; microsoft/Swin-Transformer)
112
+ 'swinmlp-t': SwinMLP_Tiny,
113
+ 'swinmlp-s': SwinMLP_Small,
114
+ 'swinmlp-b': SwinMLP_Base,
115
+ 'swinmlp-l': SwinMLP_Large,
116
+ # SwinV2 Variants
117
+ 'swinv2-t': SwinV2_Tiny, # 224 recommended; 112 works with 3 stages instead of 4
118
+ 'swinv2-s': SwinV2_Small,
119
+ 'swinv2-b': SwinV2_Base,
120
+ 'swinv2-l': SwinV2_Large,
121
+ # MobileViT (V1) Variants
122
+ 'mobilevit-xxs': MobileViT_XXS,
123
+ 'mobilevit-xs': MobileViT_XS,
124
+ 'mobilevit-s': MobileViT_S,
125
+ # MobileViTv2 Variants (separable self-attention, width multipliers)
126
+ 'mobilevitv2-0.5': MobileViTv2_050,
127
+ 'mobilevitv2-0.75': MobileViTv2_075,
128
+ 'mobilevitv2-1.0': MobileViTv2_100,
129
+ 'mobilevitv2-1.25': MobileViTv2_125,
130
+ 'mobilevitv2-1.5': MobileViTv2_150,
131
+ 'mobilevitv2-1.75': MobileViTv2_175,
132
+ 'mobilevitv2-2.0': MobileViTv2_200,
133
+ # MobileViTv3 (V1-based) Variants
134
+ 'mobilevitv3-xxs': MobileViTv3_XXS,
135
+ 'mobilevitv3-xs': MobileViTv3_XS,
136
+ 'mobilevitv3-s': MobileViTv3_S,
137
+ # MobileViTv3 (V2-based) Variants (width multipliers)
138
+ 'mobilevitv3-0.5': MobileViTv3_050,
139
+ 'mobilevitv3-0.75': MobileViTv3_075,
140
+ 'mobilevitv3-1.0': MobileViTv3_100,
141
+ 'mobilevitv3-1.25': MobileViTv3_125,
142
+ 'mobilevitv3-1.5': MobileViTv3_150,
143
+ 'mobilevitv3-1.75': MobileViTv3_175,
144
+ 'mobilevitv3-2.0': MobileViTv3_200,
145
+ }
146
+
147
+ def build_backbone(backbone_name, backbone_kwargs: Dict[str, Any], device: torch.device) -> nn.Module:
148
+ if backbone_name not in BACKBONE_REGISTRY:
149
+ raise ValueError(
150
+ f"Backbone '{backbone_name}' is not registered. "
151
+ f"Available backbones: {list(BACKBONE_REGISTRY.keys())}"
152
+ )
153
+ backbone_kwargs = copy.deepcopy(backbone_kwargs)
154
+ return BACKBONE_REGISTRY[backbone_name](**backbone_kwargs).to(device)