import sys, math, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, sweep_baseline, make_report SEEDS=tuple(range(8)); E=8; K=4; EPS=1.0 class RoutedForecaster(nn.Module): def __init__(self, win, use_expand, lam, temperature): super().__init__(); self.use_expand=use_expand; self.lam=lam self.inp=nn.Linear(1,16); self.pos=nn.Parameter(torch.zeros(1,win,16)); nn.init.normal_(self.pos,.02) self.enc=nn.TransformerEncoder(nn.TransformerEncoderLayer(16,2,32,batch_first=True,dropout=0.),1) self.router=nn.Linear(16,E) self.experts=nn.ModuleList([nn.Sequential(nn.Linear(16,16),nn.ReLU(),nn.Linear(16,16)) for _ in range(E)]) self.head=nn.Linear(16,1); self.temperature=temperature; self.last_sig={} def expansion_loss(self,p): n=p.shape[0]; rng=np.random.default_rng(1771); vals=[]; hard=[]; bounds=[] for m in [K,2*K,4*K,8*K]: if m<=n//2: for _ in range(1): ix=torch.as_tensor(rng.choice(n,m,False),device=p.device); c=1-torch.prod(1-p[ix],0); C=c.sum() b=EPS*m/(math.log(3*m/K)**2); vals.append(F.relu(torch.as_tensor(b,device=p.device)-C)**2) hard.append(int(p[ix].argmax(1).unique().numel())); bounds.append(b) return torch.stack(vals).mean(),hard,bounds def forward(self,x): h=self.enc(self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]) p=F.softmax(self.router(h)/self.temperature,-1); ys=torch.stack([e(h) for e in self.experts],2) out=self.head((ys*p.unsqueeze(-1)).sum(2)[:,-1]); load=p.mean((0,1)); lb=E*(load**2).sum() ex,hard,bounds=self.expansion_loss(p.reshape(-1,E)) counts=p.argmax(-1).reshape(-1).bincount(minlength=E).float() self.last_sig={'soft_C':float((1-torch.prod(1-p.reshape(-1,E),0)).sum().detach()),'mean_hard_neighborhood':float(np.mean(hard)),'violation_fraction':float(np.mean([anp.mean([a['mean_hard_neighborhood'] for a in base_sigs]) and np.mean([a['load_std'] for a in idea_sigs])