refactor : 清理工厂和配置系统中的死代码与冗余抽象
- 删除 Registry 中未使用的 category/priority 字段,_entries 简化为直接存储类引用 - 修正 __init_subclass__ 避免叶子类(AutoRegressiveLM 等)创建空注册表 - 删除 5 个工厂的薄 create() 覆写,统一使用 BaseFactory.create(name, *args, **kwargs) - 删除 3 处零调用的 available_types/available_strategies 别名死代码 - 删除零调用的 BaseModelConfig.to_file 死代码 - 将 BaseConfig.from_json/to_json 重命名为 from_file/to_file,消除与子类重复 - 移除两个 inference builder 中总是被覆写的 prompt_tokens=0
This commit is contained in:
+6
-10
@@ -8,8 +8,6 @@ from typing import (
|
||||
Dict,
|
||||
ForwardRef,
|
||||
Generic,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
TypeVar,
|
||||
)
|
||||
@@ -56,21 +54,19 @@ class BaseFactory(ABC, Generic[T]):
|
||||
unrelated parameters.
|
||||
"""
|
||||
|
||||
_entries: Dict[str, Tuple[Type, Optional[str], int]]
|
||||
_entries: Dict[str, Type[T]]
|
||||
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
cls._entries = {}
|
||||
for orig_base in getattr(cls, "__orig_bases__", ()):
|
||||
if _get_origin(orig_base) is BaseFactory:
|
||||
(arg,) = _get_args(orig_base)
|
||||
cls._entries = {}
|
||||
cls._component_base = _resolve_type(arg, cls)
|
||||
return
|
||||
|
||||
@classmethod
|
||||
def register(
|
||||
cls, name: str, category: Optional[str] = None, priority: int = 0
|
||||
) -> Callable[[Type[T]], Type[T]]:
|
||||
def register(cls, name: str) -> Callable[[Type[T]], Type[T]]:
|
||||
"""Decorator to register a component class.
|
||||
|
||||
Validates that the decorated class inherits from the generic
|
||||
@@ -81,7 +77,7 @@ class BaseFactory(ABC, Generic[T]):
|
||||
cls._validate_component(component_cls)
|
||||
if name in cls._entries:
|
||||
raise ValueError(f"Component '{name}' is already registered")
|
||||
cls._entries[name] = (component_cls, category, priority)
|
||||
cls._entries[name] = component_cls
|
||||
return component_cls
|
||||
|
||||
return decorator
|
||||
@@ -96,7 +92,7 @@ class BaseFactory(ABC, Generic[T]):
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
component_cls = entry[0]
|
||||
component_cls = entry
|
||||
sig = inspect.signature(component_cls.__init__)
|
||||
has_var_kwargs = any(
|
||||
p.kind == inspect.Parameter.VAR_KEYWORD for p in sig.parameters.values()
|
||||
@@ -130,7 +126,7 @@ class BaseFactory(ABC, Generic[T]):
|
||||
raise ValueError(
|
||||
f"Unknown component: '{name}'. Supported types: {sorted(cls._entries)}"
|
||||
)
|
||||
return entry[0]
|
||||
return entry
|
||||
|
||||
@classmethod
|
||||
def list_registered(cls) -> list:
|
||||
|
||||
Reference in New Issue
Block a user