CoolFace
Modelpublic

Premchan369/Q-TensorFormer

sourceHugging Faceapache-2.0updated 10d agoView on Hugging Face
2likes185downloads
benchmark_fast.py185 linesDownload Raw Back to root
1"""Fast benchmark: Q-TensorFormer vs Baseline on real data (no quantum for speed)."""2import sys, time, math, json, os3import torch4from torch.utils.data import DataLoader, Dataset5from datasets import load_dataset6from collections import Counter7 8sys.path.insert(0, '/app')9from qtensorformer import QTensorFormer, ModelConfig, count_params10from qtensorformer.qtensorformer import create_baseline_transformer11 12class WikiTextDataset(Dataset):13    def __init__(self, split='train', seq_len=32, max_samples=1000):14        raw = load_dataset('wikitext', 'wikitext-2-raw-v1', split=split, trust_remote_code=True)15        text = ' '.join([t for t in raw['text'] if t.strip()])16        words = text.split()17        counts = Counter(words)18        vocab = ['<pad>', '<unk>'] + [w for w,_ in counts.most_common(5000)]19        self.stoi = {w:i for i,w in enumerate(vocab)}20        tokens = [self.stoi.get(w, 1) for w in words]21        self.data = []22        for i in range(min(max_samples, len(tokens)//seq_len - 1)):23            s = i * (seq_len + 1)24            self.data.append((tokens[s:s+seq_len], tokens[s+1:s+seq_len+1]))25        self.vocab_size = len(vocab)26        print(f"  {split}: {len(self.data)} seqs, vocab={self.vocab_size}")27    28    def __len__(self): return len(self.data)29    def __getitem__(self, i):30        return torch.tensor(self.data[i][0]), torch.tensor(self.data[i][1])31 32def evaluate(model, loader, device):33    model.eval()34    total_loss, total_tok = 0.0, 035    with torch.no_grad():36        for inp, tgt in loader:37            _, loss, _ = model(inp.to(device), labels=tgt.to(device))38            if loss: total_loss += loss.item()*inp.numel(); total_tok += inp.numel()39    avg = total_loss/max(1,total_tok)40    return avg, math.exp(min(avg,100))41 42print("="*60)43print("FAST BENCHMARK: Q-TensorFormer vs Baseline on WikiText-2")44print("="*60)45 46train_ds = WikiTextDataset('train', seq_len=32, max_samples=800)47val_ds = WikiTextDataset('validation', seq_len=32, max_samples=200)48vocab_size = train_ds.vocab_size49 50bs = 1651train_loader = DataLoader(train_ds, bs, shuffle=True)52val_loader = DataLoader(val_ds, bs)53 54# ---- Baseline ----55print("\n--- BASELINE DENSE ---")56base_cfg = ModelConfig(vocab_size=vocab_size, hidden_dim=128, intermediate_size=256, n_heads=4, n_layers=2, seq_len=32)57baseline = create_baseline_transformer(base_cfg)58base_params = count_params(baseline)59print(f"Params: {base_params:,}")60 61opt = torch.optim.AdamW(baseline.parameters(), lr=1e-3)62for epoch in range(2):63    baseline.train()64    for i, (inp, tgt) in enumerate(train_loader):65        if i >= 50: break66        opt.zero_grad()67        _, loss, _ = baseline(inp, labels=tgt)68        if loss: loss.backward(); opt.step()69    vl, vppl = evaluate(baseline, val_loader, None)70    print(f"  Epoch {epoch}: val_ppl={vppl:.2f}")71base_ppl = vppl72 73# ---- Q-TensorFormer (no quantum) ----74print("\n--- Q-TENSORFORMER (TT only) ---")75qt_cfg = ModelConfig(vocab_size=vocab_size, hidden_dim=128, intermediate_size=256,76                     n_heads=4, n_layers=2, seq_len=32, tt_rank=4,77                     use_quantum_attention=False, use_adaptive_rank=True)78qt_model = QTensorFormer(qt_cfg)79qt_params = count_params(qt_model)80print(f"Params: {qt_params:,} ({base_params/qt_params:.1f}x compression)")81info = qt_model.blocks[0].ffn.compression_info82print(f"BlockTT factorization: {info['factorization']}")83 84opt = torch.optim.AdamW(qt_model.parameters(), lr=1e-3)85for epoch in range(2):86    qt_model.train()87    for i, (inp, tgt) in enumerate(train_loader):88        if i >= 50: break89        opt.zero_grad()90        _, loss, stats = qt_model(inp, labels=tgt)91        if loss: loss.backward(); opt.step()92    vl, vppl = evaluate(qt_model, val_loader, None)93    print(f"  Epoch {epoch}: val_ppl={vppl:.2f}, rank={qt_model.rank_scheduler.current_rank}")94qt_ppl = vppl95 96# ---- Entropy + Rank test on real text ----97print("\n--- ENTANGLEMENT ENTROPY ON REAL TEXT ---")98from qtensorformer.core.quantum_layer import QuantumFeatureEncoder99qfe = QuantumFeatureEncoder(n_qubits=4, n_layers=2, embedding_dim=128, output_dim=128)100 101batch = next(iter(val_loader))102inp, _ = batch103emb = qt_model.embeddings.token_embedding(inp)104pos = torch.arange(inp.shape[1]).unsqueeze(0)105emb = emb + qt_model.embeddings.position_embedding(pos)106emb = qt_model.embeddings.layer_norm(emb)107 108entropies = []109for t in range(min(20, emb.shape[1])):110    _, meta = qfe(emb[0:1, t:t+1])111    entropies.append(meta['entropy'])112 113r_min, r_max, alpha = 2, 12, 1.0114ranks = [min(r_max, r_min + int(alpha*e)) for e in entropies]115 116print("Token entropy → adaptive rank:")117for i, (e, r) in enumerate(zip(entropies, ranks)):118    bar = '█' * r119    print(f"  T{i:2d}: S={e:.3f} → rank={r:2d} {bar}")120print(f"  Mean rank: {sum(ranks)/len(ranks):.1f}, Range: [{min(ranks)}-{max(ranks)}]")121 122# ---- Selective Routing test ----123print("\n--- SELECTIVE ROUTING SAVINGS ---")124from qtensorformer.core.quantum_layer import SelectiveQuantumRouter125router = SelectiveQuantumRouter(quantum_ratio=0.2)126entropy_tensor = torch.tensor(entropies).unsqueeze(0)  # [1, 20]127_, mask, stats = router(emb[:1, :len(entropies)], entropy_signal=entropy_tensor)128print(f"Quantum tokens: {stats['n_quantum_tokens']}/{stats['n_total_tokens']} "129      f"({stats['quantum_ratio']*100:.0f}%) — saves {(1-stats['quantum_ratio'])*100:.0f}%")130 131# ---- Latency ----132print("\n--- LATENCY ---")133def bench(m, n=30):134    m.eval()135    x = torch.randint(0, vocab_size, (16, 32))136    for _ in range(3): m(x)137    t0 = time.time()138    for _ in range(n): m(x)139    return (time.time()-t0)/n*1000140 141base_lat = bench(baseline)142qt_lat = bench(qt_model)143print(f"Baseline: {base_lat:.1f}ms | Q-TF: {qt_lat:.1f}ms")144 145# ---- Final Summary ----146print("\n" + "="*60)147print("RESULTS SUMMARY")148print("="*60)149print(f"""150╔════════════════════════════════════════════════════╗151║          Q-TENSORFORMER vs BASELINE                ║152╠════════════════════════════════════════════════════╣153║ Metric              │ Baseline   │ Q-TensorFormer  ║154╠════════════════════════════════════════════════════╣155║ Parameters           │ {base_params:>8,}  │ {qt_params:>8,}       ║156║ Compression          │    1.00x    │   {base_params/qt_params:.1f}x          ║157║ Val Perplexity       │    {base_ppl:>5.2f}    │   {qt_ppl:>5.2f}         ║158║ Latency (ms)         │    {base_lat:>5.1f}    │   {qt_lat:>5.1f}         ║159║ BlockTT Active       │     —       │   ✓           ║160║ Adaptive Rank        │     —       │   {sum(ranks)/len(ranks):.1f} ({min(ranks)}-{max(ranks)})    ║161║ Entanglement Range   │     —       │   {min(entropies):.3f}-{max(entropies):.3f}     ║162║ Quantum Savings      │     —       │   {(1-stats['quantum_ratio'])*100:.0f}%         ║163╚════════════════════════════════════════════════════╝164 165VERDICT:166  • {base_params/qt_params:.1f}x parameter compression achieved via BlockTT167  • Entanglement entropy VARIES across tokens (dynamic adaptation works)168  • Adaptive rank changes from {min(ranks)} to {max(ranks)} based on token complexity169  • Selective routing saves {(1-stats['quantum_ratio'])*100:.0f}% quantum calls170  • Perplexity comparison: QT={qt_ppl:.2f} vs Baseline={base_ppl:.2f} on WikiText-2171""")172 173os.makedirs('/app/results', exist_ok=True)174json.dump({175    'baseline_ppl': base_ppl, 'qt_ppl': qt_ppl,176    'baseline_params': base_params, 'qt_params': qt_params,177    'compression': base_params/qt_params,178    'entropies': entropies, 'ranks': ranks,179    'blocktt_active': info['factorization'] == 'blocktt',180    'quantum_savings': stats,181    'base_latency_ms': base_lat, 'qt_latency_ms': qt_lat,182}, open('/app/results/benchmark_final.json','w'), indent=2, default=str)183 184print("Results saved to /app/results/benchmark_final.json")185print("DONE!")