CSIGv3_train_script / src /check_submission.py
XenderYang's picture
CSIGv3 AdcSR train scripts + A100 runbook
4811c23 verified
Raw History Blame Contribute Delete
2.3 kB
#!/usr/bin/env python
"""提交包检查: 命名/数量/大小/结构/可加载/确定性。
用法: python src/check_submission.py --zip xxx.zip [--names data/test_names.txt] [--test_runner]
"""
import argparse, json, os, sys, tempfile, zipfile
import torch
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--zip", required=True)
ap.add_argument("--names", default="", help="每行一个期望文件名")
ap.add_argument("--test_runner", action="store_true", help="解压并实际加载 jit 跑 1 次")
args = ap.parse_args()
size = os.path.getsize(args.zip)
checks = {"exists": True, "size_gb": round(size / 1e9, 2), "le_10gb": size <= 10e9}
with zipfile.ZipFile(args.zip) as z:
names = z.namelist()
outs = [n for n in names if "/output_dir/" in n or n.startswith("output_dir/")]
jpgs = [n for n in outs if n.lower().endswith(".jpg")]
mods = [n for n in names if "/model_dir/" in n or n.startswith("model_dir/")]
checks["output_jpg_count"] = len(jpgs)
checks["has_model"] = any(n.endswith(".pt") for n in mods)
checks["has_runner"] = any(n.endswith("runner.py") for n in mods)
if args.names:
expect = [l.strip() for l in open(args.names, encoding="utf-8") if l.strip()]
got = {os.path.basename(n) for n in jpgs}
checks["missing"] = [e for e in expect if e not in got]
checks["extra"] = sorted(got - set(expect))[:10]
if args.test_runner:
tmp = tempfile.mkdtemp()
z.extractall(tmp)
pt = [os.path.join(tmp, n) for n in mods if n.endswith(".pt")][0]
m = torch.jit.load(pt, map_location="cpu")
m.eval()
x = torch.randn(1, 3, 512, 512).half() * 0.5
with torch.no_grad():
o1 = m(x); o2 = m(x)
checks["jit_load_ok"] = True
checks["out_shape"] = list(o1.shape)
checks["deterministic_maxdiff"] = float((o1 - o2).abs().max())
for k, v in checks.items():
print(f"{k}: {v}")
bad = [k for k, v in checks.items() if v is False] or \
(checks.get("missing") if checks.get("missing") else [])
print("PASS" if not bad else f"FAIL items: {bad}")
if __name__ == "__main__":
main()