import argparse import json import os import subprocess import sys import time from pathlib import Path SCRIPT_DIR = Path(__file__).resolve().parent TRAIN_SCRIPT = SCRIPT_DIR / "_tuning.fix.chunked.py" PYTHON_BIN = "/opt/miniconda/envs/focallora/bin/python" MODELS = { "llama": { "model_path": "/data/local/hujk/models/Llama-3.1-8B-Instruct", "head_path": "/home/hujk/gitrs/Paper2026/SortedCode2/2-2_head_identification_scoring/model_score/sep_Llama-3.1-8B-Instruct/heads_sorted/all_roc_inst_0.1.json", "length_key": "llama", "max_position": 131072, }, "qwen3": { "model_path": "/data/local/hujk/models/Qwen3-8B", "head_path": "/data/local/hujk/ISH/head_scoring/qwen3-8b_sep/heads_sorted/all_roc_inst_0.1.json", "length_key": "qwen3", "max_position": 40960, }, } def run_one(args, model_name: str, data_path: str, token_len: int, batch_size: int): cfg = MODELS[model_name] out_dir = Path(args.output_dir) / model_name / f"len_{token_len:05d}_bs_{batch_size}" cmd = [ PYTHON_BIN, str(TRAIN_SCRIPT), "--model_path", cfg["model_path"], "--data_path", data_path, "--head_path", cfg["head_path"], "--output_dir", str(out_dir), "--topk", args.topk, "--epochs", "1", "--batch_size", str(batch_size), "--lr", args.lr, "--lambda_preserve", args.lambda_preserve, "--attn-chunk-size", str(args.attn_chunk_size), "--max-train-steps", "1", "--skip-save-eval", "--repeat-single-sample", str(batch_size), "--max-len", str(args.max_len), "--no-save", "--logits-to-keep", str(args.logits_to_keep), ] if args.gradient_checkpointing: cmd.append("--gradient-checkpointing") env = os.environ.copy() env["CUDA_VISIBLE_DEVICES"] = str(args.gpu) start = time.monotonic() proc = subprocess.run( cmd, cwd=str(SCRIPT_DIR.parent.parent), env=env, text=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, timeout=args.timeout, ) elapsed = time.monotonic() - start output = proc.stdout or "" tail = "\n".join(output.splitlines()[-80:]) return { "model": model_name, "data_path": data_path, "token_len": token_len, "batch_size": batch_size, "returncode": proc.returncode, "ok": proc.returncode == 0, "elapsed_sec": elapsed, "output_tail": tail, } def load_manifest(path: str): with open(path, "r", encoding="utf-8") as f: payload = json.load(f) return payload["files"] def main(): parser = argparse.ArgumentParser(description="Run one-step chunked training capacity tests.") parser.add_argument("--manifest", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-1_model_training_data_gen/len_test/manifest.json") parser.add_argument("--models", default="llama,qwen3") parser.add_argument("--gpu", default="0") parser.add_argument("--output-dir", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-2_model_training/len_test_outputs") parser.add_argument("--result-json", default="/home/hujk/gitrs/Paper2026/SortedCode2/3-2_model_training/len_test_outputs/results.json") parser.add_argument("--topk", default="100p") parser.add_argument("--lr", default="5e-4") parser.add_argument("--lambda-preserve", default="1.0") parser.add_argument("--attn-chunk-size", type=int, default=1024) parser.add_argument("--max-len", type=int, default=131072) parser.add_argument("--batch-candidates", default="1,2,4,8") parser.add_argument("--length-probe", choices=["asc", "max-first"], default="asc") parser.add_argument("--gradient-checkpointing", action="store_true") parser.add_argument("--logits-to-keep", type=int, default=1) parser.add_argument("--timeout", type=int, default=1800) args = parser.parse_args() files = load_manifest(args.manifest) model_names = [name for name in args.models.split(",") if name.strip()] batch_candidates = [int(x) for x in args.batch_candidates.split(",") if x.strip()] results = {"config": vars(args), "runs": [], "summary": {}} for model_name in model_names: cfg = MODELS[model_name] key = cfg["length_key"] length_files = [ (item["token_lengths"][key], item["path"]) for item in files if item["token_lengths"][key] <= cfg["max_position"] ] length_files.sort() last_ok = None first_fail = None if args.length_probe == "max-first": probe_files = list(reversed(length_files)) else: probe_files = length_files for token_len, path in probe_files: print(f"[len] model={model_name} len={token_len} bs=1 file={path}", flush=True) try: result = run_one(args, model_name, path, token_len, 1) except subprocess.TimeoutExpired as exc: result = { "model": model_name, "data_path": path, "token_len": token_len, "batch_size": 1, "returncode": None, "ok": False, "elapsed_sec": args.timeout, "output_tail": f"timeout after {args.timeout}s\n{exc.stdout or ''}", } results["runs"].append(result) if result["ok"]: last_ok = (token_len, path) if args.length_probe == "max-first": break else: first_fail = (token_len, path) if args.length_probe == "asc": break batch_ok = [] batch_fail = None if last_ok is not None: token_len, path = last_ok for batch_size in batch_candidates: if batch_size == 1: batch_ok.append(batch_size) continue print(f"[batch] model={model_name} len={token_len} bs={batch_size}", flush=True) try: result = run_one(args, model_name, path, token_len, batch_size) except subprocess.TimeoutExpired as exc: result = { "model": model_name, "data_path": path, "token_len": token_len, "batch_size": batch_size, "returncode": None, "ok": False, "elapsed_sec": args.timeout, "output_tail": f"timeout after {args.timeout}s\n{exc.stdout or ''}", } results["runs"].append(result) if result["ok"]: batch_ok.append(batch_size) else: batch_fail = batch_size break results["summary"][model_name] = { "max_len_bs1": last_ok[0] if last_ok else None, "max_len_file": last_ok[1] if last_ok else None, "first_failed_len_bs1": first_fail[0] if first_fail else None, "max_batch_at_max_len": max(batch_ok) if batch_ok else None, "first_failed_batch_at_max_len": batch_fail, "topk": args.topk, "attn_chunk_size": args.attn_chunk_size, } result_path = Path(args.result_json) result_path.parent.mkdir(parents=True, exist_ok=True) with result_path.open("w", encoding="utf-8") as f: json.dump(results, f, ensure_ascii=False, indent=2) print(json.dumps(results["summary"], ensure_ascii=False, indent=2)) print(f"results -> {result_path}") if __name__ == "__main__": try: main() except KeyboardInterrupt: sys.exit(130)