Spaces:
Sleeping
Sleeping
| """ | |
| 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() | |