Skip to content

Commit c40d939

Browse files
committed
fix: get_trainable_info 检查 build_model 是否为 None,避免 103 个算法全显示为可训练
- registry.py: hasattr 改为 getattr(...) is None 判断,因 ModelAlgorithmAdapter.build_model property 对所有实例都返回非 None - _torch_upgrade.py: 处理 checkpoint 格式异常时跳过 torch 推理
1 parent 9f198ae commit c40d939

2 files changed

Lines changed: 38 additions & 0 deletions

File tree

algorithms/models/deep_learning/_torch_upgrade.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -588,6 +588,11 @@ def try_torch_predict(
588588
model = getattr(algorithm, "_cached_torch_model", None)
589589
if model is None or (bvid and not getattr(algorithm, "_cached_bvid", "") == bvid):
590590
model = model_cls(**(model_kwargs or {}))
591+
if isinstance(state, (tuple, list)):
592+
state = state[0]
593+
if not isinstance(state, dict):
594+
logger.warning("[%s] checkpoint 格式异常 (type=%s),跳过 torch 推理", algo_id, type(state).__name__)
595+
return fallback_fn(video_data, threshold)
591596
model.load_state_dict(state)
592597
model.to(algorithm._device).eval()
593598
algorithm._cached_torch_model = model

algorithms/registry.py

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -256,5 +256,38 @@ def reset(cls):
256256
cls._model_adapters = {}
257257
cls._initialized = False
258258

259+
@classmethod
260+
def get_trainable_info(cls) -> List[Dict]:
261+
from algorithms.training.checkpoint_manager import CheckpointManager
262+
if not cls._initialized:
263+
cls.initialize()
264+
result = []
265+
for aid, adapter in cls._algorithms.items():
266+
build_model_fn = getattr(adapter, "build_model", None)
267+
if build_model_fn is None:
268+
continue
269+
ckpt = CheckpointManager(aid)
270+
active = ckpt.active_version()
271+
result.append({
272+
"algorithm_id": aid,
273+
"name": getattr(adapter, "name", aid),
274+
"category": getattr(adapter, "category", ""),
275+
"has_ckpt": ckpt.has_checkpoint(),
276+
"active_version": active or "",
277+
})
278+
return result
279+
280+
@classmethod
281+
def get_trainable_algorithms(cls) -> List:
282+
if not cls._initialized:
283+
cls.initialize()
284+
result = []
285+
for aid, adapter in cls._algorithms.items():
286+
build_model_fn = getattr(adapter, "build_model", None)
287+
if build_model_fn is None:
288+
continue
289+
result.append((aid, adapter.algo if hasattr(adapter, "algo") else adapter, adapter))
290+
return result
291+
259292

260293
AlgorithmRegistry.initialize()

0 commit comments

Comments
 (0)