Files
Video-Tree-TRM5/app/harness/baseline_diagnosis.py
T

241 lines
9.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""离线诊断编排:把 baseline run 的错题诊断投影为逐题信号行并断点续跑落库。
"结果驱动视频级切分"离线管线的诊断步。给定一批可诊断错题:
1. 算 remaining(跳过 store 已完成题)实现续跑幂等;
2. 对剩余错题调 core.evolution.diagnose.run_diagnosis(经 StepsJsonRunLog
包装内层 RunLog,兼容 traces 未落表的历史 run);
3. 把 error_attributions / infra / degraded 三类产物确定性投影为
DiagnosisSignalRowtier 由 split_selection.score_signal 判定);
4. 逐行 store.upsert 落盘,单行单事务 → 崩溃最多丢正在写的一行。
错误处理诚实标注(不谎称全传播):
- run_diagnosis 内部对 judge/C3 判别异常是 `except Exception`→warning→默认
lapsecore/evolution/diagnose.py:2186-2192),非全传播;judge 语义歧义
按现有保护性 lapse 处理,本编排原样接受其判定,不二次兜底。
- 网络/API 层失败经 GovernedLLMClient 重试栈后仍失败会从 run_diagnosis
向上抛出,本编排不捕获、不掩盖,直接冒泡给调用方。
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from loguru import logger
from app.harness.baseline_run_log import StepsJsonRunLog
from app.harness.split_selection import evolution_target_of, score_signal
from core.evolution.diagnose import run_diagnosis
from core.evolution.types import DiagnosisSignalRow
if TYPE_CHECKING:
from core.evolution.protocols import DiagnosisSignalStore
from core.evolution.types import DiagnosisResult
from core.types import GeneratedQuestion
@dataclass(frozen=True)
class DiagnosisDeps:
"""离线诊断编排的依赖束(一次编排的全部外部端口 + 运行参数)。
frozen 保证一次编排内依赖不可变;LLM/RunLog/SkillStore/prompts 走 Protocol
注入,便于测试替换成假实现。
属性:
run_log: 内层 RunLog 实现(提供 get_predictions/get_traces),
编排内部再用 StepsJsonRunLog 包装以兼容 traces 未落表的 run。
llm: LLM 调用端口(治理后的 GovernedLLMClient)。
skill_store: 技能文件读取端口。
prompts: 诊断模板束(DiagnosePrompts)。
tree_data: 树结构字典(多视频 {video_id: tree} 或单棵树),透传给 run_diagnosis。
concurrency: 诊断并发上限。
"""
run_log: Any
llm: Any
skill_store: Any
prompts: Any
tree_data: dict[str, Any]
concurrency: int
async def run_baseline_diagnosis(
*,
baseline_run_id: str,
diag_fingerprint: str,
wrong_ids: list[str],
questions: dict[str, GeneratedQuestion],
store: DiagnosisSignalStore,
deps: DiagnosisDeps,
) -> None:
"""对 baseline run 的错题跑离线诊断并把信号逐行落库(断点续跑幂等)。
参数:
baseline_run_id: baseline run 标识(如 "infer_adhoc"),信号行主键之一。
diag_fingerprint: 诊断口径指纹,隔离不同诊断配置的信号,主键之一。
wrong_ids: 本次待诊断的可诊断错题 question_id 列表(保序)。
questions: question_id → GeneratedQuestion 映射,需覆盖 wrong_ids 全部题
及 run_diagnosis 返回的所有 infra/degraded 题(用于取 video_id/task_type)。
store: 诊断信号存储端口,逐行 upsert 落盘并提供 done_question_ids 续跑查询。
deps: 外部依赖束(见 DiagnosisDeps)。
返回:
None。副作用为把逐题 DiagnosisSignalRow 写入 store。
关键实现:
- remaining = wrong_ids 去除 store 已完成题;空则直接 return(续跑幂等,
重复调用零副作用)。
- run_diagnosis 只诊断 remaining,避免重复 LLM 调用浪费。
- 三类产物投影互斥落库:error_attributionsdefect/lapse)、infra_question_ids
T0)、degraded_question_idsuncertain)。
"""
# Phase 1: 算 remaining(续跑幂等)
done = store.done_question_ids(baseline_run_id, diag_fingerprint)
remaining = [qid for qid in wrong_ids if qid not in done]
if not remaining:
logger.info(
"离线诊断续跑:baseline={} fingerprint={} 无剩余错题(已完成 {} 题),跳过。",
baseline_run_id,
diag_fingerprint,
len(done),
)
return
logger.info(
"离线诊断开始:baseline={} fingerprint={} 剩余 {}/{} 题待诊断。",
baseline_run_id,
diag_fingerprint,
len(remaining),
len(wrong_ids),
)
# Phase 2: 对剩余错题跑诊断(StepsJsonRunLog 兼容 traces 未落表的历史 run)
result = await run_diagnosis(
baseline_run_id,
[questions[qid] for qid in remaining],
deps.tree_data,
deps.llm,
StepsJsonRunLog(deps.run_log),
deps.skill_store,
deps.prompts,
concurrency=deps.concurrency,
question_ids=list(remaining),
only_incorrect=True,
)
# Phase 3: 投影落库
counts = _project_and_persist(
result=result,
baseline_run_id=baseline_run_id,
diag_fingerprint=diag_fingerprint,
questions=questions,
store=store,
)
logger.info(
"离线诊断落库完成:baseline={} fingerprint={} "
"T2={} T1={} T0(infra)={} uncertain(degraded)={}{} 行。",
baseline_run_id,
diag_fingerprint,
counts["T2"],
counts["T1"],
counts["T0"],
counts["uncertain"],
sum(counts.values()),
)
def _project_and_persist(
*,
result: DiagnosisResult,
baseline_run_id: str,
diag_fingerprint: str,
questions: dict[str, GeneratedQuestion],
store: DiagnosisSignalStore,
) -> dict[str, int]:
"""把 DiagnosisResult 三类产物投影为信号行并逐行 upsert,返回各 tier 计数。
参数:
result: run_diagnosis 的返回,含 error_attributions/infra/degraded 三类产物。
baseline_run_id: 信号行主键之一。
diag_fingerprint: 信号行主键之一。
questions: question_id → GeneratedQuestion,用于取 video_id/task_type。
store: 诊断信号存储端口。
返回:
{tier: 行数} 计数字典(T2/T1/T0/uncertain),供上层日志与 manifest。
关键实现:
逐行 upsert(单行单事务),中途崩溃最多丢正在写的一行;三类产物互斥,
同一 question_id 不会在两类中重复出现(run_diagnosis 保证)。
"""
counts = {"T2": 0, "T1": 0, "T0": 0, "uncertain": 0}
# error_attributionsdefect→T2 / lapse→T1 / 其它→uncertain(由 score_signal 判定)
for ea in result.error_attributions:
q = questions[ea.question_id]
tier = score_signal(cause_category=ea.cause_category, infra=False, degraded=False).tier
# error_type 是 ErrorAttribution 必填字段(永远已知),确定性派生进化目标。
evolution_target = evolution_target_of(ea.error_type)
store.upsert(
DiagnosisSignalRow(
question_id=ea.question_id,
video_id=q.video_id,
baseline_run_id=baseline_run_id,
diag_fingerprint=diag_fingerprint,
task_type=q.task_type,
error_type=ea.error_type,
cause_category=ea.cause_category,
tier=tier,
evolution_target=evolution_target,
degraded=False,
infra=False,
session_id=None,
)
)
counts[tier] = counts.get(tier, 0) + 1
# infra_question_ids:基础设施失败护栏排除 → T0,不参与训练主体
for qid in result.infra_question_ids:
q = questions[qid]
store.upsert(
DiagnosisSignalRow(
question_id=qid,
video_id=q.video_id,
baseline_run_id=baseline_run_id,
diag_fingerprint=diag_fingerprint,
task_type=q.task_type,
error_type=None,
cause_category=None,
tier="T0",
evolution_target=None,
degraded=False,
infra=True,
session_id=None,
)
)
counts["T0"] += 1
# degraded_question_idsjudge 解析失败降级 → uncertain,信号不可信排除出 T2
for qid in result.degraded_question_ids:
q = questions[qid]
store.upsert(
DiagnosisSignalRow(
question_id=qid,
video_id=q.video_id,
baseline_run_id=baseline_run_id,
diag_fingerprint=diag_fingerprint,
task_type=q.task_type,
error_type=None,
cause_category=None,
tier="uncertain",
evolution_target=None,
degraded=True,
infra=False,
session_id=None,
)
)
counts["uncertain"] += 1
return counts