CoolFace
Modelpublic

Premchan369/Q-TensorFormer

sourceHugging Faceapache-2.0updated 10d agoView on Hugging Face
2likes185downloads
scheduler.py155 linesDownload Raw Back to src
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