vishalp's picture
Deploy State-Centric Learning live demo
dbc6675 verified
Raw
History Blame Contribute Delete
21.1 kB
"""
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()