fix: count units not questions in trainability pre-flight; fail-fast when all types filtered
This commit is contained in:
+19
-8
@@ -298,37 +298,48 @@ def _filter_untrainable_types(
|
||||
eval_min_per_class: int,
|
||||
trainable_min_units: int,
|
||||
) -> tuple[Pools, list[str] | None]:
|
||||
"""剔除不可训练题型(val<eval_min_per_class 或 非test单元<trainable_min_units)。
|
||||
"""剔除不可训练题型(val 单元<eval_min_per_class 或 diag+val 单元<trainable_min_units)。
|
||||
|
||||
非test单元数 = 该题型 diag+val 题数(single 题 unit==题;等于 gate 阶梯该类候选数)。
|
||||
test 池不过滤(继续报告全题型准确率)。在 gate 建立前调用,避免样本不足的
|
||||
微型题型进入信息量阶梯导致门控崩溃。
|
||||
计数以**单元(unit)**为原子:AR pair 孪生对折叠计 1 个单元(等于 gate 阶梯该类
|
||||
候选数),非按题目计数——否则 pair 题型会以 2 倍题目数误通过阈值。test 池不过滤
|
||||
(继续报告全题型准确率)。在 gate 建立前调用,避免样本不足的微型题型进入信息量
|
||||
阶梯导致门控崩溃。
|
||||
|
||||
参数:
|
||||
pools: 冻结三池。
|
||||
task_types: 显式题型子集(None 表示全部),过滤后按 keep 收窄。
|
||||
eval_min_per_class: 验证池每类保底题数下限。
|
||||
eval_min_per_class: 验证池每类保底单元数下限。
|
||||
trainable_min_units: 每类可训练所需最小 diag+val 单元数。
|
||||
|
||||
返回:
|
||||
过滤后的 (pools, task_types):pools.diagnosis/validation 仅保留 keep 题型,
|
||||
test 原样;task_types 收窄为 keep(原 None 时返回 sorted(keep))。
|
||||
|
||||
异常:
|
||||
RuntimeError: 过滤后无任何可训练题型(切分/阈值需调整,fail-fast 不空转训练)。
|
||||
"""
|
||||
diag_by_type = Counter(q.task_type for q in pools.diagnosis)
|
||||
val_by_type = Counter(q.task_type for q in pools.validation)
|
||||
diag_by_type = Counter(u.task_type for u in build_units(pools.diagnosis))
|
||||
val_by_type = Counter(u.task_type for u in build_units(pools.validation))
|
||||
keep: set[str] = set()
|
||||
dropped: list[tuple[str, str]] = []
|
||||
for tt in set(diag_by_type) | set(val_by_type):
|
||||
n_val = val_by_type.get(tt, 0)
|
||||
n_units = diag_by_type.get(tt, 0) + n_val
|
||||
if n_val < eval_min_per_class:
|
||||
dropped.append((tt, f"val={n_val}<{eval_min_per_class}"))
|
||||
dropped.append((tt, f"val_units={n_val}<{eval_min_per_class}"))
|
||||
elif n_units < trainable_min_units:
|
||||
dropped.append((tt, f"units={n_units}<{trainable_min_units}"))
|
||||
else:
|
||||
keep.add(tt)
|
||||
for tt, why in sorted(dropped):
|
||||
logger.warning("可训练性预检剔除题型 {}({})", tt, why)
|
||||
if not keep:
|
||||
detail = ";".join(f"{tt}({why})" for tt, why in sorted(dropped))
|
||||
raise RuntimeError(
|
||||
"可训练性预检剔除了全部题型,无题型满足 "
|
||||
f"val_units>={eval_min_per_class} 且 units>={trainable_min_units}:"
|
||||
f"{detail}。请调整池切分或降低阈值。"
|
||||
)
|
||||
new_pools = replace(
|
||||
pools,
|
||||
diagnosis=[q for q in pools.diagnosis if q.task_type in keep],
|
||||
|
||||
Reference in New Issue
Block a user