owl-imdl 0.0.1__tar.gz

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.
Files changed (84) hide show
  1. owl_imdl-0.0.1/LICENSE +9 -0
  2. owl_imdl-0.0.1/PKG-INFO +79 -0
  3. owl_imdl-0.0.1/README.md +59 -0
  4. owl_imdl-0.0.1/pyproject.toml +33 -0
  5. owl_imdl-0.0.1/setup.cfg +4 -0
  6. owl_imdl-0.0.1/src/owl/__init__.py +40 -0
  7. owl_imdl-0.0.1/src/owl/_internal/__init__.py +0 -0
  8. owl_imdl-0.0.1/src/owl/_internal/fmt.py +50 -0
  9. owl_imdl-0.0.1/src/owl/_internal/lazy.py +43 -0
  10. owl_imdl-0.0.1/src/owl/_internal/logger.py +102 -0
  11. owl_imdl-0.0.1/src/owl/_internal/registry.py +66 -0
  12. owl_imdl-0.0.1/src/owl/_internal/signature.py +161 -0
  13. owl_imdl-0.0.1/src/owl/_monitor/__init__.py +0 -0
  14. owl_imdl-0.0.1/src/owl/_monitor/client.py +90 -0
  15. owl_imdl-0.0.1/src/owl/_monitor/codec.py +82 -0
  16. owl_imdl-0.0.1/src/owl/_monitor/generated/__init__.py +0 -0
  17. owl_imdl-0.0.1/src/owl/_monitor/generated/monitor_pb2.py +56 -0
  18. owl_imdl-0.0.1/src/owl/_monitor/generated/monitor_pb2_grpc.py +189 -0
  19. owl_imdl-0.0.1/src/owl/_monitor/ring.py +116 -0
  20. owl_imdl-0.0.1/src/owl/_monitor/server.py +155 -0
  21. owl_imdl-0.0.1/src/owl/_monitor/snapshot.py +66 -0
  22. owl_imdl-0.0.1/src/owl/cli/__init__.py +1 -0
  23. owl_imdl-0.0.1/src/owl/cli/cmd_init.py +64 -0
  24. owl_imdl-0.0.1/src/owl/cli/cmd_stats.py +168 -0
  25. owl_imdl-0.0.1/src/owl/cli/cmd_version.py +3 -0
  26. owl_imdl-0.0.1/src/owl/cli/main.py +68 -0
  27. owl_imdl-0.0.1/src/owl/cli/templates/finetune.py +159 -0
  28. owl_imdl-0.0.1/src/owl/cli/templates/train.py +158 -0
  29. owl_imdl-0.0.1/src/owl/cli/templates/val_metric.py +90 -0
  30. owl_imdl-0.0.1/src/owl/cli/templates/visual.py +0 -0
  31. owl_imdl-0.0.1/src/owl/engine/__init__.py +20 -0
  32. owl_imdl-0.0.1/src/owl/engine/_launch/__init__.py +21 -0
  33. owl_imdl-0.0.1/src/owl/engine/_launch/defaults.py +70 -0
  34. owl_imdl-0.0.1/src/owl/engine/_launch/dump.py +57 -0
  35. owl_imdl-0.0.1/src/owl/engine/_launch/normalize.py +282 -0
  36. owl_imdl-0.0.1/src/owl/engine/_launch/types.py +226 -0
  37. owl_imdl-0.0.1/src/owl/engine/_launch/validate.py +159 -0
  38. owl_imdl-0.0.1/src/owl/engine/app.py +420 -0
  39. owl_imdl-0.0.1/src/owl/engine/engine.py +320 -0
  40. owl_imdl-0.0.1/src/owl/engine/pipeline.py +170 -0
  41. owl_imdl-0.0.1/src/owl/engine/state.py +39 -0
  42. owl_imdl-0.0.1/src/owl/monitor/__init__.py +16 -0
  43. owl_imdl-0.0.1/src/owl/monitor/client.py +159 -0
  44. owl_imdl-0.0.1/src/owl/toolkits/__init__.py +28 -0
  45. owl_imdl-0.0.1/src/owl/toolkits/common/__init__.py +22 -0
  46. owl_imdl-0.0.1/src/owl/toolkits/common/ckpt.py +49 -0
  47. owl_imdl-0.0.1/src/owl/toolkits/common/fs.py +52 -0
  48. owl_imdl-0.0.1/src/owl/toolkits/common/image.py +55 -0
  49. owl_imdl-0.0.1/src/owl/toolkits/common/metrics.py +290 -0
  50. owl_imdl-0.0.1/src/owl/toolkits/common/seed.py +110 -0
  51. owl_imdl-0.0.1/src/owl/toolkits/common/validator.py +63 -0
  52. owl_imdl-0.0.1/src/owl/toolkits/criterion/__init__.py +25 -0
  53. owl_imdl-0.0.1/src/owl/toolkits/criterion/base.py +52 -0
  54. owl_imdl-0.0.1/src/owl/toolkits/criterion/types.py +48 -0
  55. owl_imdl-0.0.1/src/owl/toolkits/data/__init__.py +18 -0
  56. owl_imdl-0.0.1/src/owl/toolkits/data/augment/__init__.py +0 -0
  57. owl_imdl-0.0.1/src/owl/toolkits/data/augment/transforms.py +68 -0
  58. owl_imdl-0.0.1/src/owl/toolkits/data/augment/types.py +114 -0
  59. owl_imdl-0.0.1/src/owl/toolkits/data/collectors/__init__.py +0 -0
  60. owl_imdl-0.0.1/src/owl/toolkits/data/collectors/owl.py +57 -0
  61. owl_imdl-0.0.1/src/owl/toolkits/data/dataloader.py +133 -0
  62. owl_imdl-0.0.1/src/owl/toolkits/data/dataset.py +83 -0
  63. owl_imdl-0.0.1/src/owl/toolkits/data/types.py +71 -0
  64. owl_imdl-0.0.1/src/owl/toolkits/evaluator/__init__.py +27 -0
  65. owl_imdl-0.0.1/src/owl/toolkits/evaluator/base.py +20 -0
  66. owl_imdl-0.0.1/src/owl/toolkits/evaluator/default.py +83 -0
  67. owl_imdl-0.0.1/src/owl/toolkits/model/__init__.py +23 -0
  68. owl_imdl-0.0.1/src/owl/toolkits/model/base.py +76 -0
  69. owl_imdl-0.0.1/src/owl/toolkits/model/types.py +53 -0
  70. owl_imdl-0.0.1/src/owl/toolkits/optimizer/__init__.py +24 -0
  71. owl_imdl-0.0.1/src/owl/toolkits/optimizer/adamw.py +7 -0
  72. owl_imdl-0.0.1/src/owl/toolkits/optimizer/types.py +1 -0
  73. owl_imdl-0.0.1/src/owl/toolkits/scheduler/__init__.py +29 -0
  74. owl_imdl-0.0.1/src/owl/toolkits/scheduler/poly.py +13 -0
  75. owl_imdl-0.0.1/src/owl/toolkits/scheduler/types.py +1 -0
  76. owl_imdl-0.0.1/src/owl/toolkits/visualizer/__init__.py +23 -0
  77. owl_imdl-0.0.1/src/owl/toolkits/visualizer/base.py +46 -0
  78. owl_imdl-0.0.1/src/owl/toolkits/visualizer/default_mask.py +33 -0
  79. owl_imdl-0.0.1/src/owl_imdl.egg-info/PKG-INFO +79 -0
  80. owl_imdl-0.0.1/src/owl_imdl.egg-info/SOURCES.txt +82 -0
  81. owl_imdl-0.0.1/src/owl_imdl.egg-info/dependency_links.txt +1 -0
  82. owl_imdl-0.0.1/src/owl_imdl.egg-info/entry_points.txt +2 -0
  83. owl_imdl-0.0.1/src/owl_imdl.egg-info/requires.txt +7 -0
  84. owl_imdl-0.0.1/src/owl_imdl.egg-info/top_level.txt +1 -0
owl_imdl-0.0.1/LICENSE ADDED
@@ -0,0 +1,9 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2025 kwen
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy of this software and associated documentation files (the "Software"), to deal in the Software without restriction, including without limitation the rights to use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of the Software, and to permit persons to whom the Software is furnished to do so, subject to the following conditions:
6
+
7
+ The above copyright notice and this permission notice shall be included in all copies or substantial portions of the Software.
8
+
9
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE.
@@ -0,0 +1,79 @@
1
+ Metadata-Version: 2.4
2
+ Name: owl-imdl
3
+ Version: 0.0.1
4
+ Summary: IMDL template engine
5
+ Author-email: kwen <aikwen@outlook.com>
6
+ License-Expression: MIT
7
+ Classifier: Programming Language :: Python :: 3
8
+ Classifier: Operating System :: OS Independent
9
+ Requires-Python: >=3.10
10
+ Description-Content-Type: text/markdown
11
+ License-File: LICENSE
12
+ Requires-Dist: loguru>=0.7.3
13
+ Requires-Dist: rich>=13.4.0
14
+ Requires-Dist: python-statemachine>=3.0.0
15
+ Requires-Dist: questionary>=2.1.1
16
+ Requires-Dist: prettytable>=3.17.0
17
+ Requires-Dist: grpcio>=1.80.0
18
+ Requires-Dist: protobuf>=6.31.1
19
+ Dynamic: license-file
20
+
21
+ # owl-imdl
22
+ ![Python Version](https://img.shields.io/badge/python->=3.10-blue)
23
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
24
+
25
+ ## Installation (TestPyPI)
26
+
27
+ Currently, this package is available on **TestPyPI**. You can install it using the following command:
28
+
29
+ ```bash
30
+ pip install -i https://test.pypi.org/simple/ owl-imdl
31
+ ```
32
+
33
+ ## Manual Dependencies
34
+
35
+ To keep the package lightweight and flexible, this project does not enforce deep learning framework dependencies (to avoid version conflicts).
36
+
37
+ Please manually install the following dependencies according to your environment before use:
38
+
39
+ 1. **PyTorch**: Visit [pytorch](https://pytorch.org/get-started/locally/) to get the command for your CUDA version.
40
+ 2. Other Essentials:
41
+ ```bash
42
+ pip install numpy Pillow albumentations
43
+ ```
44
+
45
+ ## Quick Start
46
+
47
+ ### Initialize a Project
48
+
49
+ Run the following command in any directory to generate a training script (e.g., `my_project.py`):
50
+
51
+ ```bash
52
+ owl init
53
+ ```
54
+
55
+ The generated template uses `example/` as a placeholder dataset directory. Replace it with your own dataset path in actual use.
56
+
57
+ ## Dataset Structure
58
+
59
+ ```text
60
+ my_dataset/
61
+ ├── gt/ # Ground Truth images
62
+ ├── tp/ # Tampered/Target images
63
+ └── my_dataset.json # Index file (MUST match the folder name!)
64
+ ```
65
+
66
+ JSON Format (my_dataset.json):
67
+
68
+ ```json
69
+ [
70
+ {
71
+ "tp": "tampered_image_01.jpg",
72
+ "gt": "mask_01.png"
73
+ },
74
+ {
75
+ "tp": "tampered_image_02.jpg",
76
+ "gt": "mask_02.png"
77
+ }
78
+ ]
79
+ ```
@@ -0,0 +1,59 @@
1
+ # owl-imdl
2
+ ![Python Version](https://img.shields.io/badge/python->=3.10-blue)
3
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
4
+
5
+ ## Installation (TestPyPI)
6
+
7
+ Currently, this package is available on **TestPyPI**. You can install it using the following command:
8
+
9
+ ```bash
10
+ pip install -i https://test.pypi.org/simple/ owl-imdl
11
+ ```
12
+
13
+ ## Manual Dependencies
14
+
15
+ To keep the package lightweight and flexible, this project does not enforce deep learning framework dependencies (to avoid version conflicts).
16
+
17
+ Please manually install the following dependencies according to your environment before use:
18
+
19
+ 1. **PyTorch**: Visit [pytorch](https://pytorch.org/get-started/locally/) to get the command for your CUDA version.
20
+ 2. Other Essentials:
21
+ ```bash
22
+ pip install numpy Pillow albumentations
23
+ ```
24
+
25
+ ## Quick Start
26
+
27
+ ### Initialize a Project
28
+
29
+ Run the following command in any directory to generate a training script (e.g., `my_project.py`):
30
+
31
+ ```bash
32
+ owl init
33
+ ```
34
+
35
+ The generated template uses `example/` as a placeholder dataset directory. Replace it with your own dataset path in actual use.
36
+
37
+ ## Dataset Structure
38
+
39
+ ```text
40
+ my_dataset/
41
+ ├── gt/ # Ground Truth images
42
+ ├── tp/ # Tampered/Target images
43
+ └── my_dataset.json # Index file (MUST match the folder name!)
44
+ ```
45
+
46
+ JSON Format (my_dataset.json):
47
+
48
+ ```json
49
+ [
50
+ {
51
+ "tp": "tampered_image_01.jpg",
52
+ "gt": "mask_01.png"
53
+ },
54
+ {
55
+ "tp": "tampered_image_02.jpg",
56
+ "gt": "mask_02.png"
57
+ }
58
+ ]
59
+ ```
@@ -0,0 +1,33 @@
1
+ [build-system]
2
+ requires = ["setuptools>=61.0.0"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "owl-imdl"
7
+ version = "0.0.1"
8
+ description = "IMDL template engine"
9
+ readme = "README.md"
10
+ authors = [
11
+ {name = "kwen", email = "aikwen@outlook.com"},
12
+ ]
13
+ license = "MIT"
14
+ classifiers = [
15
+ "Programming Language :: Python :: 3",
16
+ "Operating System :: OS Independent",
17
+ ]
18
+ requires-python = ">=3.10"
19
+ dependencies = [
20
+ "loguru>=0.7.3",
21
+ "rich>=13.4.0",
22
+ "python-statemachine>=3.0.0",
23
+ "questionary>=2.1.1",
24
+ "prettytable>=3.17.0",
25
+ "grpcio>=1.80.0",
26
+ "protobuf>=6.31.1",
27
+ ]
28
+
29
+ [project.scripts]
30
+ owl = "owl.cli.main:main"
31
+
32
+ [tool.setuptools.packages.find]
33
+ where = ["src"]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1,40 @@
1
+ from typing import TYPE_CHECKING
2
+
3
+ from ._internal.lazy import attach_lazy_modules
4
+ from .toolkits.common.seed import seed_everything
5
+
6
+
7
+ try:
8
+ from importlib.metadata import PackageNotFoundError, version
9
+
10
+ try:
11
+ __version__ = version("owl-imdl")
12
+ except PackageNotFoundError:
13
+ __version__ = "unknown"
14
+ except ImportError:
15
+ __version__ = "unknown"
16
+
17
+
18
+ seed = seed_everything
19
+
20
+
21
+ # IDE 提示
22
+ if TYPE_CHECKING:
23
+ from . import engine
24
+ from . import toolkits
25
+
26
+
27
+ __all__ = attach_lazy_modules(
28
+ target_globals=globals(),
29
+ package=__package__,
30
+ delayed_modules={
31
+ "engine": ".engine",
32
+ "toolkits": ".toolkits",
33
+ },
34
+ )
35
+
36
+ __all__.extend([
37
+ "__version__",
38
+ "seed",
39
+ "seed_everything",
40
+ ])
File without changes
@@ -0,0 +1,50 @@
1
+ from prettytable import PrettyTable
2
+
3
+
4
+ def format_metrics_table(all_metrics: dict[str, dict[str, float]], current_epoch: int) -> str:
5
+ """使用 PrettyTable 生成纯 ASCII 指标表格,确保日志文件无乱码。"""
6
+ if not all_metrics:
7
+ return ""
8
+
9
+ table = PrettyTable()
10
+
11
+ # 获取指标名称 (例如: AUC, F1)
12
+ first_ds = list(all_metrics.keys())[0]
13
+ metric_keys = list(all_metrics[first_ds].keys())
14
+
15
+ # 设置表头列名
16
+ table.field_names = ["DATASET"] + [k.upper() for k in metric_keys]
17
+
18
+ for ds_name, metrics in all_metrics.items():
19
+ row_values = [ds_name]
20
+ for k in metric_keys:
21
+ val = metrics.get(k, 0.0)
22
+ # 保持 4 位小数格式化
23
+ row_values.append(f"{val:.4f}" if isinstance(val, float) else str(val))
24
+ table.add_row(row_values)
25
+
26
+ table.align = "c"
27
+
28
+ lines = [
29
+ f"\nEpoch [{current_epoch}] Summary",
30
+ table.get_string(),
31
+ ""
32
+ ]
33
+
34
+ return "\n".join(lines) + "\n"
35
+
36
+ def format_zero_pad(value: int, max_val: int) -> str:
37
+ """
38
+ 对数字进行前导补 0 格式化。
39
+ 对齐宽度由 max_val 的位数决定。
40
+
41
+ Args:
42
+ value (int): 当前值。
43
+ max_val (int): 允许达到的最大值, 用于确定宽度。
44
+
45
+ Returns:
46
+ str: 格式化后的字符串。例如:value=1, max_val=100 -> "001"
47
+ """
48
+ # 计算最大值的位数作为宽度
49
+ width = len(str(max_val))
50
+ return f"{value:0{width}d}"
@@ -0,0 +1,43 @@
1
+ from __future__ import annotations
2
+
3
+ import importlib
4
+ from types import ModuleType
5
+ from typing import Mapping
6
+
7
+
8
+ def attach_lazy_modules(
9
+ target_globals: dict,
10
+ package: str | None,
11
+ delayed_modules: Mapping[str, str],
12
+ ) -> list[str]:
13
+ """
14
+ 给当前包挂载子模块懒加载能力。
15
+
16
+ 当用户第一次访问包下的某个子模块时,才真正导入该模块。
17
+ 例如访问 ``common.fs`` 时,才导入 ``owl.toolkits.common.fs``。
18
+ Args:
19
+ target_globals (dict): 调用方模块的 globals()。
20
+ package (str | None): 调用方模块的 __package__。
21
+ delayed_modules (Mapping[str, str]): 子模块懒加载映射表。
22
+
23
+ Returns:
24
+ list[str]: 建议写入 __all__ 的模块名列表。
25
+ """
26
+
27
+ if package is None:
28
+ package = target_globals.get("__name__", "")
29
+
30
+ def __getattr__(name: str) -> ModuleType:
31
+ if name in delayed_modules:
32
+ module = importlib.import_module(delayed_modules[name], package=package)
33
+
34
+ # 缓存,避免下次再次触发 __getattr__
35
+ target_globals[name] = module
36
+
37
+ return module
38
+
39
+ current_module_name = target_globals.get("__name__", package)
40
+ raise AttributeError(f"module {current_module_name!r} has no attribute {name!r}")
41
+
42
+ target_globals["__getattr__"] = __getattr__
43
+ return list(delayed_modules.keys())
@@ -0,0 +1,102 @@
1
+ import sys
2
+ import pathlib
3
+ from loguru import logger as _logger
4
+
5
+ LOGO_TEXT = r"""
6
+ __
7
+ ____ _ __/ /
8
+ / __ \ | /| / / /
9
+ / /_/ / |/ |/ / /
10
+ \____/|__/|__/_/ owl(v{})
11
+ """
12
+
13
+ class OwlLogger:
14
+ """全局日志管理器"""
15
+ # 单例
16
+ _initialized = False
17
+
18
+ @classmethod
19
+ def setup(cls, work_dir: str | pathlib.Path):
20
+ """
21
+ 初始化日志
22
+ """
23
+ if cls._initialized:
24
+ return
25
+
26
+ work_dir = pathlib.Path(work_dir)
27
+ work_dir.mkdir(parents=True, exist_ok=True)
28
+
29
+ # 移除 loguru 默认的终端输出
30
+ _logger.remove()
31
+
32
+ # ==========================================
33
+ # 终端控制台输出 (Console)
34
+ # ==========================================
35
+ _logger.add(
36
+ sys.stdout,
37
+ format="<green>{time:YYYY-MM-DD HH:mm:ss}</green> - <level>{message}</level>",
38
+ level="INFO", # 只打印 INFO 及以上级别的日志
39
+ colorize=True, # 开启颜色
40
+ filter=lambda record: "mode" not in record["extra"],
41
+ enqueue=True # 开启异步/多线程安全
42
+ )
43
+
44
+ # ==========================================
45
+ # 训练日志文件 (train.log)
46
+ # ==========================================
47
+ train_log_path = work_dir.joinpath("train.log")
48
+ _logger.add(
49
+ str(train_log_path),
50
+ # 纯文本 format
51
+ format="{time:YYYY-MM-DD HH:mm:ss} - {message}",
52
+ level="INFO",
53
+ filter=lambda record: record["extra"].get("mode", "train") == "train", # 默认接收 train 模式的日志
54
+ rotation="100 MB", # 文件超过 100MB 自动打包
55
+ retention="30 days", # 日志保留 30 天
56
+ enqueue=True
57
+ )
58
+
59
+ # ==========================================
60
+ # 验证日志文件 (validate.log)
61
+ # ==========================================
62
+ val_log_path = work_dir.joinpath("validate.log")
63
+ _logger.add(
64
+ str(val_log_path),
65
+ format="{time:YYYY-MM-DD HH:mm:ss} - {message}",
66
+ level="INFO",
67
+ filter=lambda record: record["extra"].get("mode") == "val", # 只有被标记为 val 的日志才会存到这里
68
+ rotation="100 MB",
69
+ retention="30 days",
70
+ enqueue=True
71
+ )
72
+
73
+ # 标记为已初始化
74
+ cls._initialized = True
75
+ _logger.info(f"日志目录: {work_dir}")
76
+
77
+ @classmethod
78
+ def welcome(cls):
79
+ """打印欢迎词与 Logo"""
80
+ try:
81
+ from .. import __version__
82
+ except ImportError:
83
+ __version__ = "unknown"
84
+
85
+ logo = LOGO_TEXT.format(__version__)
86
+
87
+ _logger.opt(raw=True, colors=True).info(f"<cyan>{logo}</cyan>\n")
88
+ _logger.opt(colors=True).info("<cyan>owl engine</cyan> is starting...")
89
+
90
+ @classmethod
91
+ def stop(cls):
92
+ """打印结束语"""
93
+ if not cls._initialized:
94
+ return
95
+
96
+ _logger.opt(colors=True).info("<cyan>owl engine</cyan> is stopped!")
97
+ _logger.complete()
98
+
99
+ @classmethod
100
+ def is_initialized(cls) -> bool:
101
+ return cls._initialized
102
+ logger = _logger
@@ -0,0 +1,66 @@
1
+ import inspect
2
+ from typing import Callable, Generic, TypeVar
3
+
4
+ # 定义一个泛型类型 T
5
+ T = TypeVar('T')
6
+
7
+ class Registry(Generic[T]):
8
+ """泛型组件注册器。
9
+
10
+ 支持直接注册类(Class),也支持注册构建函数(Factory Function)。
11
+ """
12
+
13
+ def __init__(self, name: str):
14
+ self._name = name
15
+ # 存储名称到可调用对象(类或函数)的映射
16
+ self._obj_map: dict[str, Callable[..., T]] = {}
17
+
18
+ def register(self, name: str|None = None) -> Callable[[Callable[..., T]], Callable[..., T]]:
19
+ """装饰器:将类或构建函数注册到工厂中。
20
+
21
+ 把类或者函数放到 map 里面
22
+ """
23
+
24
+ def _register(obj: Callable[..., T]) -> Callable[..., T]:
25
+ key = name if name is not None else obj.__name__
26
+ if key in self._obj_map:
27
+ raise ValueError(f"组件 '{key}' 已经在注册器 '{self._name}' 中被注册过了!")
28
+
29
+ self._obj_map[key] = obj
30
+ return obj
31
+
32
+ return _register
33
+
34
+ def get(self, name: str) -> Callable[..., T]:
35
+ """获取已注册的类或构建函数"""
36
+ if name not in self._obj_map:
37
+ valid_list = list(self._obj_map.keys())
38
+ raise ValueError(f"未知的组件类型: '{name}'。在 '{self._name}' 中的支持列表: {valid_list}")
39
+ return self._obj_map[name]
40
+
41
+ def build(self, obj_type: str, **kwargs) -> T:
42
+ """组件构建
43
+
44
+ Args:
45
+ obj_type (str): 注册的组件名称(例如 "poly" 或 "adamw")。
46
+ **kwargs: 目标类或函数所需的任意参数。
47
+
48
+ Returns:
49
+ T: 实例化后的具体对象。
50
+ """
51
+ # 获取已注册的类或工厂函数
52
+ build_func = self.get(obj_type)
53
+
54
+ # kwargs 解包
55
+ try:
56
+ inspect.signature(build_func).bind(**kwargs)
57
+ except TypeError as e:
58
+ raise ValueError(f"构建组件 '{obj_type}' 失败,参数不匹配: {e}") from e
59
+
60
+ return build_func(**kwargs)
61
+
62
+ def __contains__(self, key: str) -> bool:
63
+ return key in self._obj_map
64
+
65
+ def __repr__(self) -> str:
66
+ return f"Registry(name={self._name}, items={list(self._obj_map.keys())})"
@@ -0,0 +1,161 @@
1
+ from __future__ import annotations
2
+
3
+ import inspect
4
+ from collections.abc import Callable
5
+ from typing import Any
6
+
7
+
8
+ def check_required_params(
9
+ func: Callable[..., Any],
10
+ required_params: tuple[str, ...],
11
+ ) -> None:
12
+ """检查函数签名是否包含指定的最小参数名集合。
13
+
14
+ 该函数只检查参数名,不检查类型注解。
15
+
16
+ 规则:
17
+ 1. 如果函数显式声明了 required_params 中的所有参数,则通过。
18
+ 2. 如果函数包含 **kwargs,则认为可以接收任意参数,也通过。
19
+ 3. 否则抛出 ValueError。
20
+
21
+ Args:
22
+ func: 被检查的函数或可调用对象。
23
+ required_params: 必须包含的参数名集合。
24
+
25
+ Raises:
26
+ ValueError: 当函数签名缺少 required_params 中的参数时抛出。
27
+ """
28
+ sig = inspect.signature(func)
29
+
30
+ has_var_kwargs = any(
31
+ param.kind == inspect.Parameter.VAR_KEYWORD
32
+ for param in sig.parameters.values()
33
+ )
34
+
35
+ if has_var_kwargs:
36
+ return
37
+
38
+ missing_params = [
39
+ name for name in required_params
40
+ if name not in sig.parameters
41
+ ]
42
+
43
+ if missing_params:
44
+ raise ValueError(
45
+ f"函数签名缺少必要参数: {missing_params}. "
46
+ f"required={list(required_params)}, "
47
+ f"signature={func.__name__}{sig}"
48
+ )
49
+
50
+
51
+ def _has_var_kwargs(sig: inspect.Signature) -> bool:
52
+ """判断函数签名中是否包含 **kwargs。"""
53
+ return any(
54
+ param.kind == inspect.Parameter.VAR_KEYWORD
55
+ for param in sig.parameters.values()
56
+ )
57
+
58
+
59
+ def _required_param_names_from_base(
60
+ base_sig: inspect.Signature,
61
+ ) -> tuple[str, ...]:
62
+ """从基类方法签名中提取需要子类显式兼容的参数名。
63
+
64
+ 这里只提取普通参数名,不包含 self,也不包含 *args / **kwargs。
65
+ """
66
+ names: list[str] = []
67
+
68
+ for name, param in base_sig.parameters.items():
69
+ if name == "self":
70
+ continue
71
+
72
+ if param.kind in (
73
+ inspect.Parameter.VAR_POSITIONAL,
74
+ inspect.Parameter.VAR_KEYWORD,
75
+ ):
76
+ continue
77
+
78
+ names.append(name)
79
+
80
+ return tuple(names)
81
+
82
+
83
+ def check_method_contract(
84
+ child_cls: type,
85
+ base_cls: type,
86
+ method_name: str,
87
+ *,
88
+ require_var_kwargs: bool = True,
89
+ check_return_annotation: bool = False,
90
+ ) -> None:
91
+ """检查子类方法是否满足基类方法的最小签名契约。
92
+
93
+ 该函数只检查参数名,不检查类型注解。
94
+
95
+ 规则:
96
+ 1. 子类必须具有 method_name 方法。
97
+ 2. 子类方法必须显式包含基类方法中的必要参数名。
98
+ 3. 默认要求子类方法保留 **kwargs,以兼容 Owl 后续扩展。
99
+ 4. 可选检查子类方法是否声明返回值注解。
100
+
101
+ 注意:
102
+ **kwargs 只用于接收未来扩展参数,不用于替代当前核心参数。
103
+ 例如基类 forward 要求 batch_data/current_epoch/current_step,
104
+ 那么子类也应该显式声明这些参数。
105
+
106
+ Args:
107
+ child_cls: 用户实现的子类,例如 MyModel。
108
+ base_cls: Owl 基类,例如 OwlModel。
109
+ method_name: 需要检查的方法名,例如 "forward"。
110
+ require_var_kwargs: 是否要求子类方法包含 **kwargs。
111
+ check_return_annotation: 是否要求子类方法声明返回值注解。
112
+
113
+ Raises:
114
+ TypeError: 当子类方法签名不满足基类契约时抛出。
115
+ """
116
+ if not hasattr(base_cls, method_name):
117
+ raise TypeError(
118
+ f"基类 {base_cls.__name__} 不存在方法: {method_name}"
119
+ )
120
+
121
+ if not hasattr(child_cls, method_name):
122
+ raise TypeError(
123
+ f"子类 {child_cls.__name__} 不存在方法: {method_name}"
124
+ )
125
+
126
+ base_method = getattr(base_cls, method_name)
127
+ child_method = getattr(child_cls, method_name)
128
+
129
+ base_sig = inspect.signature(base_method)
130
+ child_sig = inspect.signature(child_method)
131
+
132
+ base_required_params = _required_param_names_from_base(base_sig)
133
+ child_has_var_kwargs = _has_var_kwargs(child_sig)
134
+
135
+ missing_params = [
136
+ name for name in base_required_params
137
+ if name not in child_sig.parameters
138
+ ]
139
+
140
+ if missing_params:
141
+ raise TypeError(
142
+ f"{child_cls.__name__}.{method_name} 签名不满足 "
143
+ f"{base_cls.__name__}.{method_name} 契约。"
144
+ f"缺少参数: {missing_params}. "
145
+ f"required={list(base_required_params)}, "
146
+ f"signature={child_cls.__name__}.{method_name}{child_sig}"
147
+ )
148
+
149
+ if require_var_kwargs and not child_has_var_kwargs:
150
+ raise TypeError(
151
+ f"{child_cls.__name__}.{method_name} 缺少 **kwargs。"
152
+ f"Owl 组件方法建议保留 **kwargs,以兼容未来框架注入的扩展上下文参数。"
153
+ f"signature={child_cls.__name__}.{method_name}{child_sig}"
154
+ )
155
+
156
+ if check_return_annotation:
157
+ if child_sig.return_annotation is inspect.Signature.empty:
158
+ raise TypeError(
159
+ f"{child_cls.__name__}.{method_name} 缺少返回值类型注解。"
160
+ f"建议与 {base_cls.__name__}.{method_name} 保持一致。"
161
+ )
File without changes