Activation2 / sweep.py
w-ahmad's picture
Auto upload 2026-08-03T20:03:31.200077
425a401 verified
Raw
History Blame Contribute Delete
2.61 kB
#!/usr/bin/env python3
"""Sweep GLU activations with identical data and hyperparameters."""
import argparse
import copy
import json
import yaml
from pathlib import Path
from transformers import AutoTokenizer, set_seed
from exp import TinyLlamaConfig, TinyLlamaForCausalLM, build_dataset, create_trainer
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", required=True, help="Base YAML config")
parser.add_argument("--activations", nargs="+", default=["silu", "gelu", "relu", "mish"])
parser.add_argument("--push", action="store_true")
args = parser.parse_args()
with open(args.config) as f:
base = yaml.safe_load(f)
# Explicit seed before any randomness (dataset, samplers, model init)
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 act in args.activations:
cfg = copy.deepcopy(base)
cfg["model"]["glu_activation"] = act
out = Path(cfg["training"]["output_dir"]).parent / f"{act}_run"
cfg["training"]["output_dir"] = str(out)
cfg["training"]["hub_model_id"] = (
f"{cfg['training'].get('hub_model_id', 'tiny-llama-lab')}-{act}"
)
# Re-seed before each variant so dropout / init is identical if needed
set_seed(seed)
print(f"\n{'='*60}\n>>> Activation: {act} | Out: {out}\n{'='*60}")
config = TinyLlamaConfig(**cfg["model"])
model = TinyLlamaForCausalLM(config)
trainer = create_trainer(model, tokenizer, cfg, train_ds, eval_ds)
trainer.train()
metrics = trainer.evaluate()
results.append(
{"activation": act, "eval_loss": metrics.get("eval_loss"), "out": str(out)}
)
trainer.save_model(str(out))
if args.push or cfg["training"].get("push_to_hub", False):
trainer.push_to_hub()
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['activation']:12s} eval_loss={r['eval_loss']:.4f}")
if __name__ == "__main__":
main()