fix: make pools.json freeze atomic + add split manifest
This commit is contained in:
+19
-8
@@ -10,6 +10,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
import math
|
import math
|
||||||
|
import os
|
||||||
import random
|
import random
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
@@ -490,6 +491,22 @@ def _dict_to_q(d: dict) -> GeneratedQuestion:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _atomic_write_json(path: Path, obj: object) -> None:
|
||||||
|
"""原子写 JSON:先写 <path>.tmp 再 os.replace,避免半截文件。
|
||||||
|
|
||||||
|
崩溃或并发写入时,直接 write_text 可能留下被截断的 JSON;本助手先把完整
|
||||||
|
内容写入同目录临时文件,再用同一文件系统上的原子 rename 替换目标,
|
||||||
|
保证读者只会看到旧完整文件或新完整文件。
|
||||||
|
|
||||||
|
参数:
|
||||||
|
path: 目标 JSON 文件路径。
|
||||||
|
obj: 可 json 序列化对象。
|
||||||
|
"""
|
||||||
|
tmp = path.with_suffix(path.suffix + ".tmp")
|
||||||
|
tmp.write_text(json.dumps(obj, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||||
|
os.replace(tmp, path)
|
||||||
|
|
||||||
|
|
||||||
def save_pools(
|
def save_pools(
|
||||||
pools: Pools,
|
pools: Pools,
|
||||||
path: Path,
|
path: Path,
|
||||||
@@ -542,10 +559,7 @@ def save_pools(
|
|||||||
data["train_ratio"] = config.train_ratio
|
data["train_ratio"] = config.train_ratio
|
||||||
data["test_source"] = str(config.test_questions_dir) if config.test_questions_dir else None
|
data["test_source"] = str(config.test_questions_dir) if config.test_questions_dir else None
|
||||||
|
|
||||||
path.write_text(
|
_atomic_write_json(path, data)
|
||||||
json.dumps(data, ensure_ascii=False, indent=2),
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def load_pools(path: Path) -> Pools:
|
def load_pools(path: Path) -> Pools:
|
||||||
@@ -767,10 +781,7 @@ def build_or_load_pools(
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
# 重新冻结
|
# 重新冻结
|
||||||
pools_path.write_text(
|
_atomic_write_json(pools_path, raw)
|
||||||
json.dumps(raw, ensure_ascii=False, indent=2),
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"per_category 增量追加 {} 个新类别: {}",
|
"per_category 增量追加 {} 个新类别: {}",
|
||||||
len(new_types),
|
len(new_types),
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""结果驱动视频级切分的冻结溯源 manifest。
|
||||||
|
|
||||||
|
冻结的 pools.json 是切分产物;manifest 记录产出这份切分的关键输入
|
||||||
|
(baseline_run_id、诊断指纹、随机种子、配置)与 pools.json 的内容指纹
|
||||||
|
(pools_sha256),供后续 build_split 写溯源、以及复现校验时比对。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from app.harness.pools import _atomic_write_json
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
def write_manifest(
|
||||||
|
path: Path,
|
||||||
|
*,
|
||||||
|
baseline_run_id: str,
|
||||||
|
diag_fingerprint: str,
|
||||||
|
seed: int,
|
||||||
|
config: dict,
|
||||||
|
pools_json_text: str,
|
||||||
|
coverage_report: dict,
|
||||||
|
generated_at: str,
|
||||||
|
) -> dict:
|
||||||
|
"""写切分冻结溯源 manifest(原子写),返回写入的 dict。
|
||||||
|
|
||||||
|
pools_sha256 = sha256(pools_json_text),供复现时校验冻结的 pools.json 内容
|
||||||
|
是否与本次切分一致。generated_at 由调用方传入(库内不用 datetime.now),
|
||||||
|
以保证相同输入产出相同 manifest,可复现。
|
||||||
|
|
||||||
|
参数:
|
||||||
|
path: manifest 目标 JSON 文件路径。
|
||||||
|
baseline_run_id: 产出本次切分所依据的基线 run 标识。
|
||||||
|
diag_fingerprint: 诊断结果指纹(决定 train/val 归属的输入)。
|
||||||
|
seed: 切分使用的随机种子。
|
||||||
|
config: 切分相关配置快照(如 train_ratio 等)。
|
||||||
|
pools_json_text: 冻结的 pools.json 完整文本,用于计算内容指纹。
|
||||||
|
coverage_report: 各类别 train/val 覆盖统计报告。
|
||||||
|
generated_at: 生成时间戳(ISO 字符串),由调用方传入。
|
||||||
|
|
||||||
|
返回:
|
||||||
|
写入 manifest 的 dict(与落盘内容一致)。
|
||||||
|
"""
|
||||||
|
manifest = {
|
||||||
|
"baseline_run_id": baseline_run_id,
|
||||||
|
"diag_fingerprint": diag_fingerprint,
|
||||||
|
"seed": seed,
|
||||||
|
"config": config,
|
||||||
|
"pools_sha256": hashlib.sha256(pools_json_text.encode("utf-8")).hexdigest(),
|
||||||
|
"coverage_report": coverage_report,
|
||||||
|
"generated_at": generated_at,
|
||||||
|
}
|
||||||
|
_atomic_write_json(path, manifest)
|
||||||
|
return manifest
|
||||||
@@ -0,0 +1,63 @@
|
|||||||
|
"""pools.json 原子冻结与切分 manifest 的单元测试。
|
||||||
|
|
||||||
|
覆盖:
|
||||||
|
- _atomic_write_json 原子写:内容正确、无残留 tmp、可覆盖已有文件。
|
||||||
|
- write_manifest 溯源写:含 baseline_run_id/diag_fingerprint/seed/pools_sha256
|
||||||
|
等键,且 pools_sha256 与给定 pools.json 内容一致。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
|
||||||
|
from app.harness.pools import _atomic_write_json
|
||||||
|
from app.harness.split_manifest import write_manifest
|
||||||
|
|
||||||
|
|
||||||
|
def test_atomic_write_replaces_and_no_tmp_left(tmp_path):
|
||||||
|
p = tmp_path / "pools.json"
|
||||||
|
_atomic_write_json(p, {"a": 1})
|
||||||
|
assert json.loads(p.read_text())["a"] == 1
|
||||||
|
assert list(tmp_path.glob("*.tmp")) == [] # 无残留 tmp
|
||||||
|
|
||||||
|
|
||||||
|
def test_atomic_write_overwrites_existing(tmp_path):
|
||||||
|
p = tmp_path / "pools.json"
|
||||||
|
_atomic_write_json(p, {"a": 1})
|
||||||
|
_atomic_write_json(p, {"a": 2})
|
||||||
|
assert json.loads(p.read_text())["a"] == 2
|
||||||
|
assert list(tmp_path.glob("*.tmp")) == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_write_manifest_contains_keys_and_matching_sha256(tmp_path):
|
||||||
|
pools_json_text = json.dumps({"split_mode": "global", "test": []}, ensure_ascii=False)
|
||||||
|
manifest_path = tmp_path / "manifest.json"
|
||||||
|
|
||||||
|
result = write_manifest(
|
||||||
|
manifest_path,
|
||||||
|
baseline_run_id="run-123",
|
||||||
|
diag_fingerprint="fp-abc",
|
||||||
|
seed=42,
|
||||||
|
config={"train_ratio": 0.5},
|
||||||
|
pools_json_text=pools_json_text,
|
||||||
|
coverage_report={"visual": {"train": 3, "val": 1}},
|
||||||
|
generated_at="2026-07-15T00:00:00Z",
|
||||||
|
)
|
||||||
|
|
||||||
|
written = json.loads(manifest_path.read_text())
|
||||||
|
for key in (
|
||||||
|
"baseline_run_id",
|
||||||
|
"diag_fingerprint",
|
||||||
|
"seed",
|
||||||
|
"config",
|
||||||
|
"pools_sha256",
|
||||||
|
"coverage_report",
|
||||||
|
"generated_at",
|
||||||
|
):
|
||||||
|
assert key in written
|
||||||
|
|
||||||
|
expected_sha = hashlib.sha256(pools_json_text.encode("utf-8")).hexdigest()
|
||||||
|
assert written["pools_sha256"] == expected_sha
|
||||||
|
assert result == written # 返回值与落盘内容一致
|
||||||
|
assert list(tmp_path.glob("*.tmp")) == []
|
||||||
Reference in New Issue
Block a user