Download src/check_submission.py from XenderYang/CSIGv3_train_script: direct link, hf CLI and curl.
- Browser
- Download file 2.3 kB
-
https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/check_submission.py
- Command line
-
hf download hf://XenderYang/CSIGv3_train_script/src/check_submission.py
-
curl -L -o check_submission.py https://huggingface.co/XenderYang/CSIGv3_train_script/resolve/main/src/check_submission.py
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() | |