""" End-to-end experiment runner across all tokenizer × domain × model × mode. Usage: uv run python -m code.experiments.run_all [--tokenizers wl simhash random] [--domains blocks] This script orchestrates: 1. Embedding generation (if not already done) via generate_multi_embeddings 2. Model training via train_lstm / train_xgb 3. Inference via inference_lstm / inference_xgb 4. Results collection and table generation """ import argparse import glob import json import os import subprocess import sys from tqdm import tqdm from code.experiments.config import ( DOMAINS, MODEL_CONFIGS, SEED, SPLITS_EVAL, TOKENIZATION_CONFIGS, ) from code.experiments.results import ExperimentTracker TOKENIZER_ALIASES = { "graphs": "wl", } def canonical_tokenizer_name(name: str) -> str: """Normalize tokenizer aliases to canonical experiment config keys.""" return TOKENIZER_ALIASES.get(name, name) def get_tokenizer_config(tokenizer: str) -> tuple[str, dict]: """Resolve tokenizer alias and return canonical name + config.""" canonical = canonical_tokenizer_name(tokenizer) if canonical not in TOKENIZATION_CONFIGS: valid = sorted(set(TOKENIZATION_CONFIGS.keys()) | set(TOKENIZER_ALIASES.keys())) raise ValueError( f"Unknown tokenizer '{tokenizer}'. " f"Valid tokenizers: {', '.join(valid)}" ) return canonical, TOKENIZATION_CONFIGS[canonical] def run_command(cmd: list[str], desc: str = "") -> int: """Run a subprocess and return exit code.""" print(f"\n>>> {desc}") print(f" {' '.join(cmd)}") result = subprocess.run(cmd, capture_output=False) if result.returncode != 0: print(f" [WARN] Command returned non-zero exit code: {result.returncode}") return result.returncode def build_wandb_cli_args( enabled: bool, project: str, entity: str | None, group: str | None, mode: str, run_name: str, tags: list[str], ) -> list[str]: """Build reusable W&B CLI arguments.""" if not enabled: return [] args = [ "--wandb", "--wandb_project", project, "--wandb_mode", mode, "--wandb_run_name", run_name, ] if entity: args.extend(["--wandb_entity", entity]) if group: args.extend(["--wandb_group", group]) if tags: args.extend(["--wandb_tags", ",".join(tags)]) return args def generate_embeddings(tokenizer: str, domain: str, data_dir: str = "data") -> None: """Generate embeddings for a tokenizer-domain pair if not already done.""" canonical_tokenizer, tok_config = get_tokenizer_config(tokenizer) enc_dir = tok_config["encoding_dir"] output_dir = os.path.join(data_dir, "encodings", enc_dir) check_dir = os.path.join(output_dir, domain, "train") # Skip if already generated if os.path.exists(check_dir) and len(os.listdir(check_dir)) > 0: print(f" [Skip] Embeddings already exist: {check_dir}") return cmd = [ sys.executable, "-m", "code.encoding_generation.generate_multi_embeddings", "--tokenizer", canonical_tokenizer, "--domain", domain, "--data_dir", data_dir, "--output_dir", output_dir, "--model_dir", os.path.join(data_dir, "encodings", "models"), ] # Add tokenizer-specific params params = tok_config["params"] for key, value in params.items(): cmd.extend([f"--{key}", str(value)]) run_command(cmd, f"Generating {tokenizer} embeddings for {domain}") def train_model( model_type: str, mode: str, tokenizer: str, domain: str, device: str = "auto", num_workers: int = 8, lstm_amp: bool = True, fast: bool = False, xgb_n_jobs: int = 8, wandb: bool = False, wandb_project: str = "state-centric-plan", wandb_entity: str | None = None, wandb_group: str | None = None, wandb_mode: str = "online", wandb_tags: str = "", data_dir: str = "data", checkpoint_dir: str = "checkpoints", ) -> str: """Train a model and return the save directory.""" _, tok_config = get_tokenizer_config(tokenizer) enc_dir = tok_config["encoding_dir"] data_path = os.path.join(data_dir, "encodings", enc_dir) save_dir = os.path.join(checkpoint_dir, enc_dir, f"{model_type}_{mode}") model_config = MODEL_CONFIGS[model_type][f"{mode}_mode"] run_tags = [tokenizer, domain, model_type, mode, "train"] extra_tags = [t.strip() for t in wandb_tags.split(",") if t.strip()] run_tags.extend(extra_tags) run_name = f"{tokenizer}-{domain}-{model_type}-{mode}-train" if model_type == "lstm": cmd = [ sys.executable, "-m", "code.modeling.train_lstm", "--domain", domain, "--data_dir", data_path, "--save_dir", save_dir, "--epochs", str(model_config["epochs"]), "--batch_size", str(model_config["batch_size"]), "--hidden_dim", str(model_config["hidden_dim"]), "--lr", str(model_config["lr"]), "--device", device, "--num_workers", str(num_workers), "--seed", str(SEED), ] if mode == "delta": cmd.append("--delta") if model_config.get("no_projection"): cmd.append("--no_projection") if lstm_amp: cmd.append("--amp") else: cmd.append("--no_amp") if fast: cmd.append("--fast") cmd.extend( build_wandb_cli_args( enabled=wandb, project=wandb_project, entity=wandb_entity, group=wandb_group, mode=wandb_mode, run_name=run_name, tags=run_tags, ) ) elif model_type == "xgboost": xgb_device = device if device in {"auto", "cuda", "cpu"} else "auto" cmd = [ sys.executable, "-m", "code.modeling.train_xgb", "--domain", domain, "--data_dir", data_path, "--save_dir", save_dir, "--encoding", enc_dir, "--n_estimators", str(model_config["n_estimators"]), "--max_depth", str(model_config["max_depth"]), "--lr", str(model_config["lr"]), "--early_stopping", str(model_config["early_stopping"]), "--device", xgb_device, "--n_jobs", str(xgb_n_jobs), "--seed", str(SEED), ] if mode == "delta": cmd.append("--delta") cmd.extend( build_wandb_cli_args( enabled=wandb, project=wandb_project, entity=wandb_entity, group=wandb_group, mode=wandb_mode, run_name=run_name, tags=run_tags, ) ) run_command(cmd, f"Training {model_type}/{mode} on {tokenizer}/{domain}") return save_dir def run_inference( model_type: str, mode: str, tokenizer: str, domain: str, device: str = "auto", lstm_amp: bool = True, fast: bool = False, xgb_n_jobs: int = 8, wandb: bool = False, wandb_project: str = "state-centric-plan", wandb_entity: str | None = None, wandb_group: str | None = None, wandb_mode: str = "online", wandb_tags: str = "", data_dir: str = "data", checkpoint_dir: str = "checkpoints", results_dir: str = "results", val_path: str | None = None, execute: bool = True, ) -> dict: """Run inference for a trained model and return split metrics.""" _, tok_config = get_tokenizer_config(tokenizer) enc_dir = tok_config["encoding_dir"] model_dir = os.path.join(checkpoint_dir, enc_dir, f"{model_type}_{mode}") output_dir = os.path.join(results_dir, enc_dir, f"{model_type}_{mode}") os.makedirs(output_dir, exist_ok=True) model_config = MODEL_CONFIGS[model_type][f"{mode}_mode"] run_tags = [tokenizer, domain, model_type, mode, "inference"] extra_tags = [t.strip() for t in wandb_tags.split(",") if t.strip()] run_tags.extend(extra_tags) run_name = f"{tokenizer}-{domain}-{model_type}-{mode}-inference" if model_type == "lstm": checkpoint_path = os.path.join(model_dir, f"{domain}_lstm_best.pt") cmd = [ sys.executable, "-m", "code.modeling.inference_lstm", "--domain", domain, "--checkpoint", checkpoint_path, "--data_dir", data_dir, "--results_dir", output_dir, "--encoding", enc_dir, "--pddl_dir", os.path.join(data_dir, "pddl"), "--device", device, "--hidden_dim", str(model_config["hidden_dim"]), "--tag", mode, "--seed", str(SEED), ] if mode == "delta": cmd.append("--delta") if model_config.get("no_projection"): cmd.append("--no_projection") if lstm_amp: cmd.append("--amp") else: cmd.append("--no_amp") if fast: cmd.append("--fast") if val_path: cmd.extend(["--val_path", val_path]) cmd.extend( build_wandb_cli_args( enabled=wandb, project=wandb_project, entity=wandb_entity, group=wandb_group, mode=wandb_mode, run_name=run_name, tags=run_tags, ) ) elif model_type == "xgboost": xgb_device = device if device in {"auto", "cuda", "cpu"} else "auto" cmd = [ sys.executable, "-m", "code.modeling.inference_xgb", "--domain", domain, "--checkpoint_dir", model_dir, "--data_dir", data_dir, "--results_dir", output_dir, "--pddl_dir", os.path.join(data_dir, "pddl"), "--device", xgb_device, "--n_jobs", str(xgb_n_jobs), "--tag", mode, "--seed", str(SEED), ] if mode == "delta": cmd.append("--delta") if val_path: cmd.extend(["--val_path", val_path]) cmd.extend( build_wandb_cli_args( enabled=wandb, project=wandb_project, entity=wandb_entity, group=wandb_group, mode=wandb_mode, run_name=run_name, tags=run_tags, ) ) if execute: run_command(cmd, f"Inference {model_type}/{mode} on {tokenizer}/{domain}") metrics = {} for split in SPLITS_EVAL: pattern = os.path.join(output_dir, f"{domain}_*_{split}_{mode}_results.json") matches = glob.glob(pattern) if not matches: # Fallback for legacy filenames without tag suffix. legacy_pattern = os.path.join(output_dir, f"{domain}_*_{split}_results.json") matches = glob.glob(legacy_pattern) if not matches: continue result_file = max(matches, key=os.path.getmtime) with open(result_file, "r") as f: rows = json.load(f) total = len(rows) solved = sum(1 for row in rows if row.get("solved")) executable = sum(1 for row in rows if row.get("val_executable")) metrics[split] = { "solved_rate": (solved / total) if total else 0.0, "exec_rate": (executable / total) if total else 0.0, } return metrics def main(): parser = argparse.ArgumentParser(description="Run full experiment suite.") parser.add_argument( "--tokenizers", nargs="+", default=list(TOKENIZATION_CONFIGS.keys()), help="Tokenizers to evaluate", ) parser.add_argument("--domains", nargs="+", default=DOMAINS) parser.add_argument( "--models", nargs="+", default=["lstm", "xgboost"], help="Model types", ) parser.add_argument("--modes", nargs="+", default=["state", "delta"]) parser.add_argument("--data_dir", default="data") parser.add_argument("--checkpoint_dir", default="checkpoints") parser.add_argument("--results_dir", default="results") parser.add_argument( "--device", choices=["auto", "cuda", "mps", "cpu"], default="auto", help="Preferred compute device policy", ) parser.add_argument("--num_workers", type=int, default=8) parser.add_argument("--xgb_n_jobs", type=int, default=8) parser.add_argument( "--lstm_amp", dest="lstm_amp", action="store_true", help="Enable CUDA mixed precision for LSTM train/inference", ) parser.add_argument( "--no_lstm_amp", dest="lstm_amp", action="store_false", help="Disable CUDA mixed precision for LSTM train/inference", ) parser.add_argument( "--fast", action="store_true", help="Enable fast CUDA settings in LSTM components", ) parser.add_argument( "--lstm_epochs", type=int, default=None, help="Optional override for LSTM epochs in both state and delta modes.", ) parser.add_argument( "--val_path", default=os.environ.get("VAL_PATH"), help="Optional path to VAL binary. If unset, local defaults are auto-detected.", ) parser.add_argument( "--skip_embedding", action="store_true", help="Skip embedding generation (assume existing)", ) parser.add_argument( "--skip_training", action="store_true", help="Skip training (assume existing models)", ) parser.add_argument( "--skip_inference", action="store_true", help="Skip inference (just collect results)", ) parser.add_argument( "--wandb", action="store_true", help="Enable W&B logging for train/inference subprocesses", ) parser.add_argument( "--wandb_project", default="state-centric-plan", help="W&B project name", ) parser.add_argument( "--wandb_entity", default=None, help="Optional W&B entity/team", ) parser.add_argument( "--wandb_group", default=None, help="Optional W&B run group for the full sweep", ) parser.add_argument( "--wandb_mode", choices=["online", "offline", "disabled"], default="online", help="W&B mode", ) parser.add_argument( "--wandb_tags", default="", help="Optional comma-separated additional W&B tags for all runs", ) parser.set_defaults(lstm_amp=True) args = parser.parse_args() alias_notes = [] for tok in args.tokenizers: canonical, _ = get_tokenizer_config(tok) if tok != canonical: alias_notes.append(f"{tok}->{canonical}") if not args.val_path: local_val_candidates = [ os.path.join("VAL", "build", "bin", "Validate.exe"), os.path.join("VAL", "build", "bin", "Validate"), os.path.join("VAL", "bin", "Validate.exe"), os.path.join("VAL", "bin", "Validate"), ] for candidate in local_val_candidates: if os.path.exists(candidate): args.val_path = candidate break if not args.skip_inference: if args.val_path: print(f"Using VAL: {args.val_path}") else: print( "[WARN] VAL path not provided/found. Inference will run, but solved/executable " "validation may fail depending on environment defaults." ) if args.lstm_epochs is not None: MODEL_CONFIGS["lstm"]["state_mode"]["epochs"] = args.lstm_epochs MODEL_CONFIGS["lstm"]["delta_mode"]["epochs"] = args.lstm_epochs print( f"Overriding LSTM epochs to {args.lstm_epochs} " f"for both state and delta modes." ) tracker = ExperimentTracker(output_dir=args.results_dir) total = ( len(args.tokenizers) * len(args.domains) * len(args.models) * len(args.modes) ) print(f"Total configurations: {total}") print(f"Tokenizers: {args.tokenizers}") if alias_notes: print(f"Tokenizer aliases resolved: {', '.join(alias_notes)}") print(f"Domains: {args.domains}") print(f"Models: {args.models}") print(f"Modes: {args.modes}") print( f"Device policy: {args.device} | LSTM AMP: {args.lstm_amp} | " f"LSTM workers: {args.num_workers} | XGB n_jobs: {args.xgb_n_jobs}" ) print(f"W&B enabled: {args.wandb} | mode: {args.wandb_mode}") print(f"Progress view: one bar per configuration ({total} total for this sweep).") done = 0 overall_bar = tqdm( total=total, desc="Overall Configurations", unit="cfg", dynamic_ncols=True, ) for tokenizer in args.tokenizers: for domain in args.domains: # Step 1: Generate embeddings if not args.skip_embedding: generate_embeddings(tokenizer, domain, args.data_dir) for model in args.models: for mode in args.modes: done += 1 print(f"\n{'='*60}") print( f"[{done}/{total}] " f"{tokenizer}/{domain}/{model}/{mode}" ) print(f"{'='*60}") config_bar = tqdm( total=3, desc=f"{done:02d}/{total} {tokenizer}/{domain}/{model}/{mode}", unit="step", dynamic_ncols=True, leave=True, ) # Step 2: Train config_bar.set_postfix_str("training") if not args.skip_training: train_model( model_type=model, mode=mode, tokenizer=tokenizer, domain=domain, device=args.device, num_workers=args.num_workers, lstm_amp=args.lstm_amp, fast=args.fast, xgb_n_jobs=args.xgb_n_jobs, wandb=args.wandb, wandb_project=args.wandb_project, wandb_entity=args.wandb_entity, wandb_group=args.wandb_group, wandb_mode=args.wandb_mode, wandb_tags=args.wandb_tags, data_dir=args.data_dir, checkpoint_dir=args.checkpoint_dir, ) config_bar.update(1) # Step 3: Inference / Result Collection config_bar.set_postfix_str("inference") metrics_by_split = run_inference( model, mode, tokenizer, domain, args.device, args.lstm_amp, args.fast, args.xgb_n_jobs, args.wandb, args.wandb_project, args.wandb_entity, args.wandb_group, args.wandb_mode, args.wandb_tags, args.data_dir, args.checkpoint_dir, args.results_dir, args.val_path, execute=not args.skip_inference, ) config_bar.update(1) config_bar.set_postfix_str("logging") for split, metrics in metrics_by_split.items(): tracker.log_result( domain=domain, tokenizer=tokenizer, model=model, mode=mode, split=split, metrics=metrics, ) config_bar.update(1) config_bar.set_postfix_str("done") config_bar.close() overall_bar.update(1) overall_bar.close() # Step 4: Generate comparison tables print(f"\n{'='*60}") print("Generating comparison tables...") print(f"{'='*60}") tracker.generate_comparison_table() tracker.save_results() print("\nDone!") if __name__ == "__main__": main()