Premchan369/Q-TensorFormer
2185
1"""2Adaptive TT-Rank Scheduler.3 4Core novelty of Q-TensorFormer: adjusts tensor rank dynamically5based on per-input complexity, estimated via attention entropy.6 7r(input) = r_min + α × normalized_entropy × (r_max - r_min)8 9Supports:10 - EMA smoothing to prevent oscillation11 - Budget-capped ranks12 - Deterministic rounding with hysteresis13"""14 15import torch16import torch.nn as nn17import math18 19 20class RankScheduler(nn.Module):21 """22 Attention entropy → TT-rank scheduler.23 24 Parameters25 ----------26 r_min : int27 Minimum tensor rank (maximum compression).28 r_max : int29 Maximum tensor rank (minimum compression).30 alpha : float31 Sensitivity: how much entropy changes the rank.32 alpha=0 → fixed rank r_min.33 alpha=1 → rank fully spans r_min to r_max.34 alpha=2.0 → aggressive scaling (default).35 smoothing : float36 EMA decay factor (0.9 = smooth, 0 = no history).37 """38 39 def __init__(self, r_min: int = 2, r_max: int = 8,40 alpha: float = 2.0, smoothing: float = 0.9):41 super().__init__()42 self.r_min = r_min43 self.r_max = r_max44 self.alpha = alpha45 self.smoothing = smoothing46 47 self.register_buffer("_ema_entropy", torch.tensor(0.5))48 self.register_buffer("_ema_rank", torch.tensor((r_min + r_max) // 2, dtype=torch.float))49 self.register_buffer("_counter", torch.tensor(0, dtype=torch.long))50 51 # Optionally learn alpha52 self.learned_alpha = nn.Parameter(torch.tensor(float(alpha)), requires_grad=False)53 54 def forward(self, entropy: torch.Tensor, seq_len: int = None) -> int:55 """56 Compute rank from attention entropy.57 58 Args:59 entropy: Scalar or 0-dim tensor (mean attention entropy).60 seq_len: Sequence length for normalization (optional).61 62 Returns:63 Integer tensor rank.64 """65 if entropy.dim() > 0:66 entropy = entropy.mean()67 68 # Normalize entropy to [0, 1]69 if seq_len is not None and seq_len > 1:70 norm_factor = math.log(seq_len)71 normalized = torch.clamp(entropy / max(norm_factor, 1e-8), 0.0, 1.0)72 else:73 normalized = torch.clamp(torch.tanh(entropy / 2.0), 0.0, 1.0)74 75 # EMA smoothing76 self._ema_entropy.mul_(self.smoothing).add_(normalized, alpha=1.0 - self.smoothing)77 smoothed = self._ema_entropy78 79 # Map to rank: r = r_min + alpha * norm * (r_max - r_min)80 alpha_val = self.learned_alpha.item()81 span = self.r_max - self.r_min82 raw = self.r_min + alpha_val * smoothed.item() * span83 84 # Round with hysteresis85 self._ema_rank.mul_(0.7).add_(raw, alpha=0.3)86 rank = int(torch.round(self._ema_rank).item())87 88 # Clamp + counter89 rank = max(self.r_min, min(self.r_max, rank))90 self._counter.add_(1)91 return rank92 93 def reset(self):94 """Reset EMA state."""95 self._ema_entropy.fill_(0.5)96 self._ema_rank.fill_((self.r_min + self.r_max) / 2.0)97 self._counter.fill_(0)98 99 @property100 def current_rank(self) -> float:101 return self._ema_rank.item()102 103 @property104 def current_entropy(self) -> float:105 return self._ema_entropy.item()106 107 108class BudgetAwareScheduler(nn.Module):109 """110 Extends RankScheduler with deployment budget constraints.111 112 Automatically caps tensor rank to meet:113 - Max parameter budget114 - Max latency target115 - Max energy per query116 """117 118 def __init__(self, scheduler: RankScheduler,119 max_params: int = None,120 max_latency_ms: float = None,121 max_energy_uj: float = None):122 super().__init__()123 self.scheduler = scheduler124 self.max_params = max_params125 self.max_latency_ms = max_latency_ms126 self.max_energy_uj = max_energy_uj127 128 def forward(self, entropy: torch.Tensor, seq_len: int = None,129 param_factors: dict = None) -> int:130 """131 Compute rank with budget constraints.132 133 Args:134 entropy: Attention entropy.135 seq_len: Sequence length.136 param_factors: Dict mapping rank → estimated total parameters.137 138 Returns:139 Budget-constrained rank.140 """141 rank = self.scheduler(entropy, seq_len)142 143 if param_factors and self.max_params:144 # Find highest rank that meets budget145 while rank > self.scheduler.r_min:146 est = param_factors.get(rank, float("inf"))147 if est <= self.max_params:148 break149 rank -= 1150 151 return rank152 153 def reset(self):154 self.scheduler.reset()155 