import json, math, random, time import numpy as np import torch import torch.nn as nn import torch.nn.functional as F SEED = 605 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) def tt_svd_matrix(W, modes, rel_tol=0.0, max_rank=None): """TT-SVD of W with C-order input/output mode ordering.""" m, n = modes d = len(m) T = W.reshape(*(list(m) + list(n))) # Convert to interleaved physical ordering (m1,n1,m2,n2,...). perm = [k for pair in zip(range(d), range(d, 2*d)) for k in pair] T = T.transpose(perm) norm = np.linalg.norm(W) eps2 = (rel_tol * norm) ** 2 # Equal per-bond budget gives the standard global error bound. budget = eps2 / max(1, d - 1) cores, ranks, tails = [], [1], [] left = T rprev = 1 for k in range(d - 1): left = left.reshape(rprev * m[k] * n[k], -1) u, s, vh = np.linalg.svd(left, full_matrices=False) tail = np.cumsum(s[::-1] ** 2)[::-1] rank = len(s) if rel_tol > 0: valid = np.where(np.r_[tail[1:], 0.0] <= budget + 1e-14)[0] if len(valid): rank = int(valid[0] + 1) if max_rank is not None: rank = min(rank, max_rank) rank = max(1, rank) discarded = float(np.sum(s[rank:] ** 2)) tails.append(discarded) cores.append(u[:, :rank].reshape(rprev, m[k], n[k], rank)) left = (s[:rank, None] * vh[:rank]) rprev = rank; ranks.append(rank) cores.append(left.reshape(rprev, m[-1], n[-1], 1)) ranks.append(1) return cores, ranks, tails, norm def tt_reconstruct(cores): x = cores[0] for c in cores[1:]: x = np.tensordot(x, c, axes=([-1], [0])) # x: m1,n1,m2,n2,..., with singleton bond ends removed d = len(cores) x = np.squeeze(x, axis=(0, -1)) shape = x.shape out_modes = [shape[2*k] for k in range(d)] in_modes = [shape[2*k+1] for k in range(d)] perm = list(range(0, 2*d, 2)) + list(range(1, 2*d, 2)) return x.transpose(perm).reshape(int(np.prod(out_modes)), int(np.prod(in_modes))) def tt_param_count(modes, ranks): m, n = modes return sum(ranks[k] * m[k] * n[k] * ranks[k+1] for k in range(len(m))) class TTLinear(nn.Module): def __init__(self, W, modes, rel_tol=0.0, max_rank=None, bias=True): super().__init__(); self.modes = modes; self.max_rank = max_rank cores, ranks, _, _ = tt_svd_matrix(W, modes, rel_tol, max_rank) self.cores = nn.ParameterList([nn.Parameter(torch.tensor(c, dtype=torch.float32)) for c in cores]) self.ranks = ranks self.bias = nn.Parameter(torch.zeros(W.shape[0])) if bias else None def forward(self, x): # Direct tensor-network contraction: batch, all input modes, TT cores. d = len(self.modes[0]) letters = list("abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ") out = letters[1:1+d]; inn = letters[1+d:1+2*d] bond = letters[1+2*d:1+3*d+1] # includes both singleton boundary bonds xsub = letters[0] + ''.join(inn) terms = [xsub] for k, c in enumerate(self.cores): terms.append(bond[k] + out[k] + inn[k] + bond[k+1]) equation = ','.join(terms) + '->' + letters[0] + ''.join(out) operands = [x.reshape(x.shape[0], *self.modes[1])] + list(self.cores) y = torch.einsum(equation, *operands).reshape(x.shape[0], -1) return y + self.bias if self.bias is not None else y @torch.no_grad() def round(self, rel_tol): W = tt_reconstruct([c.detach().cpu().numpy() for c in self.cores]) cores, ranks, _, _ = tt_svd_matrix(W, self.modes, rel_tol, self.max_rank) self.cores = nn.ParameterList([nn.Parameter(torch.tensor(c, device=self.cores[0].device, dtype=self.cores[0].dtype)) for c in cores]) self.ranks = ranks def verify(): modes = ([4,4,4], [4,4,4]); W = np.random.randn(64,64) cores, ranks, tails, norm = tt_svd_matrix(W, modes, rel_tol=0.08, max_rank=8) Wa = tt_reconstruct(cores); err = np.linalg.norm(W-Wa) # Uncapped TT-SVD tests the stated tolerance theorem directly. exact_cores, exact_ranks, exact_tails, _ = tt_svd_matrix(W, modes, rel_tol=0.08, max_rank=None) exact_err = np.linalg.norm(W - tt_reconstruct(exact_cores)) bound = 0.08 * norm # contraction equality check on a random vector layer = TTLinear(W, modes, rel_tol=0.08, max_rank=8, bias=False) x = torch.randn(3,64) direct = layer(x).detach().numpy(); expected = x.numpy() @ Wa.T return {'relative_fro_error': float(err/norm), 'requested_bound': float(bound/norm), 'bound_holds_with_cap': bool(err <= bound + 1e-6), 'uncapped_relative_fro_error': float(exact_err/norm), 'uncapped_bound_holds': bool(exact_err <= bound + 1e-6), 'contraction_max_abs': float(np.max(np.abs(direct-expected))), 'ranks': ranks, 'uncapped_ranks': exact_ranks, 'dense_params': 4096, 'tt_params': tt_param_count(modes, ranks), 'uncapped_tt_params': tt_param_count(modes, exact_ranks), 'tail_squared_sum': float(sum(tails))} def train_compare(): device = 'cuda' if torch.cuda.is_available() else 'cpu' try: torch.manual_seed(SEED) modes = ([4,4,4], [4,4,4]); dim=64 # Teacher is genuinely TT-low-rank, making this a representation test. teacher = TTLinear(np.random.randn(dim,dim)*0.25, modes, rel_tol=0, max_rank=2, bias=True).to(device) with torch.no_grad(): teacher.bias.normal_(0, .1) X = torch.randn(1024, dim, device=device); Y = torch.tanh(teacher(X)).detach() results = {} for name, model in [('dense', nn.Sequential(nn.Linear(dim,dim), nn.Tanh()).to(device)), ('tt', None)]: if model is None: model = nn.Sequential(TTLinear(np.random.randn(dim,dim)*.05, modes, rel_tol=0, max_rank=4).to(device), nn.Tanh()) opt = torch.optim.Adam(model.parameters(), lr=3e-3) t0=time.perf_counter(); losses=[] for step in range(250): pred=model(X); loss=F.mse_loss(pred,Y); opt.zero_grad(); loss.backward(); opt.step() if name=='tt' and (step+1)%50==0: model[0].round(0.03) losses.append(float(loss.detach().cpu())) results[name]={'final_mse':losses[-1], 'mse_step_50':losses[49], 'seconds':time.perf_counter()-t0, 'params':sum(p.numel() for p in model.parameters()), 'ranks': getattr(model[0], 'ranks', None) if name=='tt' else None} return {'device':device, 'results':results} except Exception as e: if device == 'cuda': torch.cuda.empty_cache(); return train_compare_cpu() raise def train_compare_cpu(): torch.cuda.is_available=lambda: False return train_compare() if __name__ == '__main__': print(json.dumps({'verification':verify(), 'training':train_compare()}, indent=2))