Activ / sweep.py
w-ahmad's picture
Auto upload 2026-08-06T20:45:13.807790
d782871 verified
Raw
History Blame Contribute Delete
5.34 kB
#!/usr/bin/env python3
"""Sweep explicit GLU / MLP variants with identical data and hyperparameters."""
import argparse
import copy
import json
import re
import time
import os
from pathlib import Path
import yaml
import wandb
import torch
from transformers import AutoTokenizer, set_seed
from exp import TinyLlamaConfig, TinyLlamaForCausalLM, build_dataset, create_trainer
def format_param_count(total_params: int) -> str:
"""Return human‑readable string with M or B suffix, 1 decimal."""
if total_params >= 1e9:
return f"{total_params / 1e9:.1f}B"
else:
return f"{total_params / 1e6:.1f}M"
def parse_variant(variant: str):
"""
Parse variant string into (prefix, activation, layers).
Formats:
glu-silu -> ('glu', 'silu', None)
mlp-s10-10L -> ('mlp', 's10', 10)
glu-relu-8L -> ('glu', 'relu', 8)
"""
# Pattern: optional layers at end with 'L'
match = re.fullmatch(r'(glu|mlp)-([a-zA-Z0-9]+)(?:-(\d+)L)?', variant)
if not match:
raise ValueError(
f"Invalid variant format: '{variant}'. "
"Expected: <glu|mlp>-<activation>[-<layers>L] e.g. glu-silu-10L"
)
prefix, act, layers_str = match.groups()
layers = int(layers_str) if layers_str is not None else None
return prefix, act, layers
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True, help="Base YAML config")
parser.add_argument(
"--variants",
nargs="+",
required=True,
help="List of variants: e.g. glu-silu-10L mlp-relu-8L"
)
parser.add_argument("--push", action="store_true")
args = parser.parse_args()
with open(args.config) as f:
base = yaml.safe_load(f)
seed = base.get("training", {}).get("seed", 42)
set_seed(seed)
tok_name = base["model"].get("tokenizer_name", "meta-llama/Llama-2-7b-hf")
tokenizer = AutoTokenizer.from_pretrained(tok_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
msl = base["model"].get("max_position_embeddings", 512)
train_ds = build_dataset(tokenizer, max_seq_len=msl, split="train")
eval_ds = build_dataset(tokenizer, max_seq_len=msl, split="validation")
results = []
for variant in args.variants:
prefix, act, layers = parse_variant(variant)
# Validation: mlp-situglu is banned
if prefix == "mlp" and act == "situglu":
raise ValueError(
f"Activation 'situglu' requires a gated architecture (GLU). "
f"Please use 'glu-situglu' instead."
)
# Build config overrides
cfg = copy.deepcopy(base)
cfg["model"]["mlp_type"] = prefix
cfg["model"]["activation"] = act
if layers is not None:
cfg["model"]["num_hidden_layers"] = layers
# Create a descriptive suffix for folders / run names
layer_suffix = f"-{layers}L" if layers is not None else ""
variant_label = f"{prefix}-{act}{layer_suffix}"
# Unique output directory
out_dir = Path(cfg["training"]["output_dir"]).parent / f"{variant_label}_run"
cfg["training"]["output_dir"] = str(out_dir)
# Re-seed for reproducibility across variants
set_seed(seed)
print(f"\n{'='*60}\n>>> Variant: {variant_label} | Out: {out_dir}\n{'='*60}")
# Instantiate model to get parameter count BEFORE trainer
config = TinyLlamaConfig(**cfg["model"])
# Use SDPA backend (or flash_attention_2 if you have it)
config._attn_implementation = "sdpa" # consistent across runs
model = TinyLlamaForCausalLM(config)
model = model.to(torch.bfloat16)
total_params = sum(p.numel() for p in model.parameters())
param_str = format_param_count(total_params)
timestamp = time.strftime("%Y%m%d-%H%M%S")
run_name = f"LM-{variant_label}-{param_str}-{timestamp}"
cfg["training"]["run_name"] = run_name
# Also update hub_model_id to include variant and layers
hub_id_base = cfg["training"].get("hub_model_id", "tiny-llama-lab")
cfg["training"]["hub_model_id"] = f"{hub_id_base}-{variant_label}"
# Ensure a fresh WandB run – remove any global WANDB_RUN_ID
os.environ.pop("WANDB_RUN_ID", None)
trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds)
try:
trainer.train()
metrics = trainer.evaluate()
results.append({
"variant": variant_label,
"eval_loss": metrics.get("eval_loss"),
"out": str(out_dir),
"run_name": run_name,
})
trainer.save_model(str(out_dir))
if args.push or cfg["training"].get("push_to_hub", False):
trainer.push_to_hub()
finally:
# Explicitly finish WandB run to avoid re‑using the same run
wandb.finish()
# Save summary
summary = Path(base["training"]["output_dir"]).parent / "sweep_summary.json"
summary.write_text(json.dumps(results, indent=2))
print("\nSweep complete:")
for r in results:
print(f" {r['variant']:20s} eval_loss={r['eval_loss']:.4f}")
if __name__ == "__main__":
main()