""" Experiment Runner: Multi-Objective Pareto Frontier Analysis. Computes Pareto curves: - Quality vs Latency - Quality vs Memory - Quality vs Energy Across presets: QTF_FULL, QTF_BALANCED, QTF_LATENCY, QTF_MEMORY, QTF_ENERGY, QTF_EDGE, QTF_CLASSICAL_ONLY. """ import sys import os import json import argparse from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.parent)) import torch from src.config import ModelConfig from src.models import QTensorFormer, DenseBaseline from src.hardware_cost_model import HardwareCostModel def main(): parser = argparse.ArgumentParser() parser.add_argument("--output", type=str, default="outputs/pareto_results.json") args = parser.parse_args() os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True) print("=" * 65) print("EXPERIMENT: Multi-Objective Pareto Frontier Generation") print("=" * 65) hw = HardwareCostModel() cfg = ModelConfig(vocab_size=1000, d_model=128, n_layers=2, tt_rank=8, use_quantum=True) presets = [ ("QTF_FULL", "Quality-First"), ("QTF_BALANCED", "Balanced"), ("QTF_LATENCY", "Latency-First"), ("QTF_MEMORY", "Memory-First"), ("QTF_ENERGY", "Energy-First"), ("QTF_EDGE", "Edge Constrained"), ("QTF_CLASSICAL_ONLY", "Classical Only"), ] points = [] for preset_code, label in presets: model = QTensorFormer(cfg, preset=preset_code) meas = hw.profile_execution(model, batch_size=1, seq_len=32, n_repeats=5) pt = { "preset": preset_code, "label": label, "latency_ms": meas.latency_ms, "peak_memory_mb": meas.peak_memory_mb, "traffic_bytes_per_token": meas.memory_traffic_bytes_per_token, "joules_per_token": meas.joules_per_token, "active_params": model.active_params, "classification": "MEASURED", } points.append(pt) print(f"{label:<18} | Lat: {meas.latency_ms:>5.2f}ms | Mem: {meas.peak_memory_mb:>5.2f}MB | J/tok: {meas.joules_per_token*1e6:>5.2f}uJ") # Add Dense baseline for reference dense = DenseBaseline(cfg) dense_meas = hw.profile_execution(dense, batch_size=1, seq_len=32, n_repeats=5) points.append({ "preset": "DENSE_BASELINE", "label": "Dense Baseline", "latency_ms": dense_meas.latency_ms, "peak_memory_mb": dense_meas.peak_memory_mb, "traffic_bytes_per_token": dense_meas.memory_traffic_bytes_per_token, "joules_per_token": dense_meas.joules_per_token, "active_params": dense.total_params, "classification": "MEASURED", }) with open(args.output, "w") as f: json.dump(points, f, indent=2) print(f"\nPareto data saved to {args.output}") if __name__ == "__main__": main()