""" HybridTimeScaleLM – Optimized Linear Model Architecture (model_linear.py) ======================================================================== Optimized alternative to `model.py` featuring Chunked Parallel Linear Attention. - VRAM Growth: Strictly linear O(S) scaling with sequence length. - Speed & Parallelism: Fully vectorized GPU kernel math, maintaining parallel execution speed. - Context Length: Designed for long token lengths (2048, 4096, 8192, 16384+). - Weight & State-Dict Compatibility: 100% drop-in replacement for `model.py` checkpoints. """ import math import os import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.checkpoint import checkpoint from transformers import ( AutoConfig, AutoModelForCausalLM, GenerationMixin, PretrainedConfig, PreTrainedModel, ) from transformers.modeling_outputs import CausalLMOutputWithPast, ModelOutput from dataclasses import dataclass from typing import Optional, Tuple, List @dataclass(init=False) class HybridTimeScaleOutput(ModelOutput): """ Base class for model's outputs that also contains a past key/values. """ loss: Optional[torch.FloatTensor] = None logits: torch.FloatTensor = None past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None hidden_states: Optional[Tuple[torch.FloatTensor]] = None attentions: Optional[Tuple[torch.FloatTensor]] = None last_hidden_state: Optional[torch.FloatTensor] = None # Global default tokenizer ID GLOBAL_TOKENIZER_ID = "mistralai/Mistral-7B-v0.3" # Prevent protobuf/sentencepiece version conflicts when AutoTokenizer loads Mistral/Llama tokenizers on macOS os.environ.setdefault("PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION", "python") # Ensure `transformers` automatically strips `_orig_mod.` prefixes left by `torch.compile` / `Trainer` wrappers # (Note: _orig_mod. prefixes should be stripped from the safetensors file directly before uploading, not patched at runtime). # ────────────────────────────────────────────────────────────────────── # Config # ────────────────────────────────────────────────────────────────────── class HybridTimeScaleConfig(PretrainedConfig): model_type = "hybrid_timescale_lm" def __init__( self, vocab_size=50304, latent_dim=768, num_layers=12, num_modes=64, layer_types=None, time_scale=128.0, dropout=0.05, pad_token_id=0, bos_token_id=1, eos_token_id=2, tie_word_embeddings=True, chunk_size=128, **kwargs, ): self.vocab_size = vocab_size self.latent_dim = latent_dim self.num_layers = num_layers self.num_modes = num_modes self.time_scale = time_scale self.dropout = dropout self.chunk_size = chunk_size if layer_types is None: layer_types = [ "softmax" if (i % 4 == 3) else "linear" for i in range(num_layers) ] assert len(layer_types) == num_layers, ( f"layer_types length ({len(layer_types)}) must equal num_layers ({num_layers})" ) self.layer_types = layer_types super().__init__( pad_token_id=pad_token_id, bos_token_id=bos_token_id, eos_token_id=eos_token_id, tie_word_embeddings=tie_word_embeddings, **kwargs, ) # ────────────────────────────────────────────────────────────────────── # Optimized Chunked Parallel Linear Fourier Mixer # ────────────────────────────────────────────────────────────────────── class LinearFourierMixer(nn.Module): """ Optimized Linear Fourier Mixer using Chunked Parallel Linear Attention. Computes intra-chunk parallel attention and inter-chunk cumulative state scan without allocating large O(S^2) attention matrices or quadratic memory. """ def __init__(self, channels, num_modes=64, num_heads=12, time_scale=128, dropout=0.05, chunk_size=128): super().__init__() assert channels % num_heads == 0, ( f"channels ({channels}) must be perfectly divisible by num_heads ({num_heads})" ) self.channels = channels self.num_modes = num_modes self.num_heads = num_heads self.head_dim = channels // num_heads self.time_scale = time_scale self.chunk_size = chunk_size freq_bands = torch.exp(torch.linspace(math.log(0.0001), math.log(num_modes), num_modes)) self.num_modes = freq_bands.shape[0] self.register_buffer("frequencies", freq_bands) self.q_proj = nn.Linear(channels, self.num_heads * self.num_modes) self.k_proj = nn.Linear(channels, self.num_heads * self.num_modes) self.v_proj = nn.Linear(channels, channels) self.proj_v2 = nn.Linear(channels, channels) self.out_proj = nn.Linear(channels, channels) self.activation = nn.SiLU() self.norm_in = nn.LayerNorm(channels) self.norm_out = nn.LayerNorm(channels) self.dropout = nn.Dropout(dropout) def forward(self, x, attention_mask=None, position_ids=None, past_key_value=None): B, seq_len, C = x.shape norm_x = self.norm_in(x) Q = F.elu(self.q_proj(norm_x)).view(B, seq_len, self.num_heads, self.num_modes) + 1.0 K = F.elu(self.k_proj(norm_x)).view(B, seq_len, self.num_heads, self.num_modes) + 1.0 v1 = self.v_proj(norm_x) v2 = self.activation(self.proj_v2(norm_x)) if position_ids is None: position_ids = torch.arange(seq_len, device=x.device, dtype=torch.long).unsqueeze(0) t = (position_ids.unsqueeze(-1).to(dtype=x.dtype) / self.time_scale) omega_t = 2 * math.pi * t * self.frequencies U = torch.cos(omega_t).unsqueeze(2) # [B, S, 1, M] V = torch.sin(omega_t).unsqueeze(2) # [B, S, 1, M] Q_cos = Q * U Q_sin = Q * V K_cos = K * U K_sin = K * V Q_rot = torch.cat([Q_cos, Q_sin], dim=-1) # [B, seq_len, H, 2M] K_rot = torch.cat([K_cos, K_sin], dim=-1) # [B, seq_len, H, 2M] v1_heads = v1.view(B, seq_len, self.num_heads, self.head_dim) if attention_mask is not None: if attention_mask.shape[1] > seq_len: mask = attention_mask[:, -seq_len:].unsqueeze(-1).unsqueeze(-1).to(dtype=x.dtype) else: mask = attention_mask.unsqueeze(-1).unsqueeze(-1).to(dtype=x.dtype) K_rot = K_rot * mask v1_heads = v1_heads * mask scale = 1.0 / math.sqrt(self.num_modes * 2) orig_dtype = Q_rot.dtype Q_rot_f = Q_rot.view(B, seq_len, self.num_heads, 2 * self.num_modes).transpose(1, 2).float() K_rot_f = K_rot.view(B, seq_len, self.num_heads, 2 * self.num_modes).transpose(1, 2).float() v1_f = v1_heads.view(B, seq_len, self.num_heads, self.head_dim).transpose(1, 2).float() if past_key_value is not None: # Recurrent O(1) step cum_kv_past, cum_k_past = past_key_value curr_kv = torch.matmul(K_rot_f.transpose(-1, -2), v1_f) curr_k = K_rot_f.sum(dim=-2, keepdim=True).transpose(-1, -2) # [B, H, 2M, 1] cum_kv_new = cum_kv_past + curr_kv cum_k_new = cum_k_past + curr_k num_total = torch.matmul(Q_rot_f, cum_kv_new) * scale denom_total = torch.matmul(Q_rot_f, cum_k_new) * scale denom_total = denom_total.clamp(min=1e-4) v1_token_mixed = (num_total / denom_total) v1_token_mixed = torch.clamp(v1_token_mixed, min=-100.0, max=100.0) v1_token_mixed = v1_token_mixed.to(orig_dtype).transpose(1, 2).reshape(B, seq_len, self.channels) present_key_value = (cum_kv_new, cum_k_new) else: # Full sequence parallel chunking (Prefill) chunk_size = self.chunk_size pad_len = (chunk_size - (seq_len % chunk_size)) % chunk_size if pad_len > 0: Q_rot = F.pad(Q_rot, (0, 0, 0, 0, 0, pad_len)) K_rot = F.pad(K_rot, (0, 0, 0, 0, 0, pad_len)) v1_heads = F.pad(v1_heads, (0, 0, 0, 0, 0, pad_len)) S_padded = seq_len + pad_len N_chunks = S_padded // chunk_size Q_c = Q_rot.view(B, N_chunks, chunk_size, self.num_heads, 2 * self.num_modes).transpose(2, 3) K_c = K_rot.view(B, N_chunks, chunk_size, self.num_heads, 2 * self.num_modes).transpose(2, 3) V_c = v1_heads.view(B, N_chunks, chunk_size, self.num_heads, self.head_dim).transpose(2, 3) Q_c_f = Q_c.float() K_c_f = K_c.float() V_c_f = V_c.float() A_intra = torch.matmul(Q_c_f, K_c_f.transpose(-1, -2)) * scale causal_mask = torch.tril(torch.ones(chunk_size, chunk_size, device=x.device, dtype=torch.float32)) A_intra = A_intra * causal_mask.unsqueeze(0).unsqueeze(0).unsqueeze(0) num_intra = torch.matmul(A_intra, V_c_f) denom_intra = A_intra.sum(dim=-1, keepdim=True) chunk_kv = torch.matmul(K_c_f.transpose(-1, -2), V_c_f) chunk_kv_past = torch.cat([torch.zeros_like(chunk_kv[:, :1]), chunk_kv[:, :-1]], dim=1) cum_kv_past = torch.cumsum(chunk_kv_past, dim=1) chunk_k_sum = K_c_f.sum(dim=-2, keepdim=True).transpose(-1, -2) chunk_k_past = torch.cat([torch.zeros_like(chunk_k_sum[:, :1]), chunk_k_sum[:, :-1]], dim=1) cum_k_past = torch.cumsum(chunk_k_past, dim=1) num_inter = torch.matmul(Q_c_f, cum_kv_past) * scale denom_inter = torch.matmul(Q_c_f, cum_k_past) * scale num_total = num_intra + num_inter denom_total = (denom_intra + denom_inter).clamp(min=1e-4) if num_total.requires_grad: num_total.register_hook(lambda grad: torch.clamp(grad, min=-30000.0, max=30000.0)) if denom_total.requires_grad: denom_total.register_hook(lambda grad: torch.clamp(grad, min=-30000.0, max=30000.0)) v1_token_mixed = (num_total / denom_total) v1_token_mixed = torch.clamp(v1_token_mixed, min=-100.0, max=100.0) v1_token_mixed = v1_token_mixed.to(orig_dtype).transpose(2, 3).reshape(B, S_padded, self.channels) if pad_len > 0: v1_token_mixed = v1_token_mixed[:, :seq_len] # Compute final state for the cache cum_kv_final = cum_kv_past[:, -1] + chunk_kv[:, -1] cum_k_final = cum_k_past[:, -1] + chunk_k_sum[:, -1] present_key_value = (cum_kv_final, cum_k_final) v1_token_mixed = self.dropout(v1_token_mixed) if attention_mask is not None: v1_token_mixed = torch.nan_to_num(v1_token_mixed, nan=0.0, posinf=0.0, neginf=0.0) if attention_mask.shape[1] > seq_len: mask = attention_mask[:, -seq_len:].unsqueeze(-1).to(dtype=v1_token_mixed.dtype) else: mask = attention_mask.unsqueeze(-1).to(dtype=v1_token_mixed.dtype) v1_token_mixed = v1_token_mixed * mask v3 = v1_token_mixed * v2 return self.norm_out(self.out_proj(v3)) + x, present_key_value # ────────────────────────────────────────────────────────────────────── # Softmax Fourier Mixer # ────────────────────────────────────────────────────────────────────── class SoftmaxFourierMixer(nn.Module): def __init__(self, channels, num_modes=64, num_heads=12, time_scale=128.0, dropout=0.05): super().__init__() assert channels % num_heads == 0, ( f"channels ({channels}) must be perfectly divisible by num_heads ({num_heads})" ) self.channels = channels self.num_modes = num_modes self.num_heads = num_heads self.head_dim = channels // num_heads self.time_scale = time_scale freq_bands = torch.exp(torch.linspace(math.log(0.0001), math.log(num_modes), num_modes)) self.num_modes = freq_bands.shape[0] self.register_buffer("frequencies", freq_bands) self.q_proj = nn.Linear(channels, self.num_heads * self.num_modes) self.k_proj = nn.Linear(channels, self.num_heads * self.num_modes) self.v_proj = nn.Linear(channels, channels) self.proj_v2 = nn.Linear(channels, channels) self.out_proj = nn.Linear(channels, channels) self.activation = nn.SiLU() self.norm_in = nn.LayerNorm(channels) self.norm_out = nn.LayerNorm(channels) self.dropout = nn.Dropout(dropout) def forward(self, x, attention_mask=None, position_ids=None, past_key_value=None): B, seq_len, C = x.shape norm_x = self.norm_in(x) Q = self.q_proj(norm_x).view(B, seq_len, self.num_heads, self.num_modes) K = self.k_proj(norm_x).view(B, seq_len, self.num_heads, self.num_modes) v1 = self.v_proj(norm_x) v2 = self.activation(self.proj_v2(norm_x)) if position_ids is None: position_ids = torch.arange(seq_len, device=x.device, dtype=torch.long).unsqueeze(0) t = (position_ids.unsqueeze(-1).to(dtype=x.dtype) / self.time_scale) omega_t = 2 * math.pi * t * self.frequencies U = torch.cos(omega_t).unsqueeze(2) V = torch.sin(omega_t).unsqueeze(2) Q_cos = Q * U Q_sin = Q * V K_cos = K * U K_sin = K * V Q_rot = torch.cat([Q_cos, Q_sin], dim=-1) K_rot = torch.cat([K_cos, K_sin], dim=-1) v1_heads = v1.view(B, seq_len, self.num_heads, self.head_dim) Q_b = Q_rot.transpose(1, 2) K_b = K_rot.transpose(1, 2) V_b = v1_heads.transpose(1, 2) if past_key_value is not None: K_past, V_past = past_key_value K_b = torch.cat([K_past, K_b], dim=2) V_b = torch.cat([V_past, V_b], dim=2) present_key_value = (K_b, V_b) seq_len_kv = K_b.size(2) if x.device.type == "mps" or (seq_len > 512 and x.device.type != "cuda"): scale = 1.0 / math.sqrt(Q_b.size(-1)) if seq_len > 256: out_chunks = [] chunk_size = 256 if seq_len > 1024 else 512 for i_start in range(0, seq_len, chunk_size): i_end = min(i_start + chunk_size, seq_len) Q_chunk = Q_b[:, :, i_start:i_end, :] # Causal chunking math for long sequences (typically prefill) K_past_chunk = K_b[:, :, :i_end + (seq_len_kv - seq_len), :] V_past_chunk = V_b[:, :, :i_end + (seq_len_kv - seq_len), :] scores_chunk = torch.matmul(Q_chunk, K_past_chunk.transpose(-2, -1)) * scale i_abs = torch.arange(i_start, i_end, device=x.device).view(-1, 1) + (seq_len_kv - seq_len) j_abs = torch.arange(i_end + (seq_len_kv - seq_len), device=x.device).view(1, -1) causal_mask = (j_abs <= i_abs) scores_chunk = scores_chunk.masked_fill(~causal_mask.unsqueeze(0).unsqueeze(0), float("-inf")) if attention_mask is not None: pad_mask = attention_mask[:, None, None, :i_end + (seq_len_kv - seq_len)].to(dtype=torch.bool) scores_chunk = scores_chunk.masked_fill(~pad_mask, float("-inf")) attn_weights = F.softmax(scores_chunk, dim=-1) out_chunk = torch.matmul(attn_weights, V_past_chunk) out_chunks.append(out_chunk) v1_token_mixed = torch.cat(out_chunks, dim=2) else: scale = 1.0 / math.sqrt(Q_b.size(-1)) scores = torch.matmul(Q_b, K_b.transpose(-2, -1)) * scale i_abs = torch.arange(seq_len, device=x.device).view(-1, 1) + (seq_len_kv - seq_len) j_abs = torch.arange(seq_len_kv, device=x.device).view(1, -1) causal_mask = (j_abs <= i_abs) scores = scores.masked_fill(~causal_mask.unsqueeze(0).unsqueeze(0), float("-inf")) if attention_mask is not None: pad_mask = attention_mask[:, None, None, :].to(dtype=torch.bool) scores = scores.masked_fill(~pad_mask, float("-inf")) attn_weights = F.softmax(scores, dim=-1) v1_token_mixed = torch.matmul(attn_weights, V_b) else: attn_mask = None if attention_mask is not None: attn_mask = attention_mask[:, None, None, :].to(dtype=Q_b.dtype) attn_mask = (1.0 - attn_mask) * torch.finfo(Q_b.dtype).min # If past_key_value is present, seq_len=1 so causality isn't needed. is_causal = past_key_value is None try: v1_token_mixed = F.scaled_dot_product_attention( Q_b, K_b, V_b, attn_mask=attn_mask, is_causal=is_causal, ) except Exception: scale = 1.0 / math.sqrt(Q_b.size(-1)) scores = torch.matmul(Q_b, K_b.transpose(-2, -1)) * scale i_abs = torch.arange(seq_len, device=x.device).view(-1, 1) + (seq_len_kv - seq_len) j_abs = torch.arange(seq_len_kv, device=x.device).view(1, -1) causal_mask = (j_abs <= i_abs) scores = scores.masked_fill(~causal_mask.unsqueeze(0).unsqueeze(0), float("-inf")) if attention_mask is not None: pad_mask = attention_mask[:, None, None, :].to(dtype=torch.bool) scores = scores.masked_fill(~pad_mask, float("-inf")) attn_weights = F.softmax(scores, dim=-1) v1_token_mixed = torch.matmul(attn_weights, V_b) v1_token_mixed = v1_token_mixed.transpose(1, 2).reshape(B, seq_len, C) v1_token_mixed = self.dropout(v1_token_mixed) if attention_mask is not None: v1_token_mixed = torch.nan_to_num(v1_token_mixed, nan=0.0, posinf=0.0, neginf=0.0) if attention_mask.shape[1] > seq_len: mask = attention_mask[:, -seq_len:].unsqueeze(-1).to(dtype=v1_token_mixed.dtype) else: mask = attention_mask.unsqueeze(-1).to(dtype=v1_token_mixed.dtype) v1_token_mixed = v1_token_mixed * mask v3 = v1_token_mixed * v2 out = self.norm_out(self.out_proj(v3)) + x return out, present_key_value # ────────────────────────────────────────────────────────────────────── # Transformer block # ────────────────────────────────────────────────────────────────────── class HybridSpectralBlock(nn.Module): def __init__(self, latent_dim, num_modes=64, is_softmax=False, time_scale=128.0, dropout=0.05, num_heads=None, chunk_size=128): super().__init__() self.is_softmax = is_softmax num_heads = num_heads if num_heads is not None else max(1, latent_dim // 64) if is_softmax: self.mixer = SoftmaxFourierMixer(latent_dim, num_modes, num_heads, time_scale, dropout) else: self.mixer = LinearFourierMixer(latent_dim, num_modes, num_heads, time_scale, dropout, chunk_size=chunk_size) self.ffn = nn.Sequential( nn.LayerNorm(latent_dim), nn.Linear(latent_dim, 4 * latent_dim), nn.GELU(), nn.Linear(4 * latent_dim, latent_dim), nn.Dropout(dropout), ) self.gradient_checkpointing = False def forward(self, x, attention_mask=None, position_ids=None, past_key_value=None): z, present_key_value = self.mixer(x, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value) out = z + self.ffn(z) return out, present_key_value # ────────────────────────────────────────────────────────────────────── # Full model # ────────────────────────────────────────────────────────────────────── class HybridTimeScalePreTrainedModel(PreTrainedModel): config_class = HybridTimeScaleConfig base_model_prefix = "hybrid_timescale" supports_gradient_checkpointing = True _no_split_modules = ["HybridSpectralBlock"] _tied_weights_keys = {"lm_head.weight": "embedding.weight"} _supports_loss_kwargs = False _supports_cache_class = False def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) elif isinstance(module, nn.LayerNorm): torch.nn.init.zeros_(module.bias) torch.nn.init.ones_(module.weight) class HybridTimeScaleLM(HybridTimeScalePreTrainedModel, GenerationMixin): def __init__(self, config): super().__init__(config) self.config = config self.embedding = nn.Embedding(config.vocab_size, config.latent_dim, padding_idx=config.pad_token_id) chunk_size = getattr(config, "chunk_size", 128) blocks = [] for layer_type in config.layer_types: blocks.append(HybridSpectralBlock( config.latent_dim, config.num_modes, is_softmax=(layer_type == "softmax"), time_scale=config.time_scale, dropout=config.dropout, chunk_size=chunk_size, )) self.mixers = nn.ModuleList(blocks) self.ln_f = nn.LayerNorm(config.latent_dim) self.lm_head = nn.Linear(config.latent_dim, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.embedding def set_input_embeddings(self, value): self.embedding = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embedding): self.lm_head = new_embedding def forward( self, input_ids=None, attention_mask=None, position_ids=None, past_key_values=None, inputs_embeds=None, labels=None, use_cache=None, output_attentions=None, output_hidden_states=None, return_dict=None, **kwargs, ): output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) return_dict = return_dict if return_dict is not None else self.config.use_return_dict use_cache = use_cache if use_cache is not None else getattr(self.config, "use_cache", True) if input_ids is not None and inputs_embeds is not None: raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time") elif input_ids is not None: batch_size, seq_length = input_ids.shape elif inputs_embeds is not None: batch_size, seq_length, _ = inputs_embeds.shape else: raise ValueError("You have to specify either input_ids or inputs_embeds") if inputs_embeds is None: hidden_states = self.embedding(input_ids) else: hidden_states = inputs_embeds if position_ids is None: device = input_ids.device if input_ids is not None else inputs_embeds.device position_ids = torch.arange(seq_length, dtype=torch.long, device=device) position_ids = position_ids.unsqueeze(0).expand(batch_size, -1) all_hidden_states = () if output_hidden_states else None presents = () if use_cache else None for i, mixer in enumerate(self.mixers): if output_hidden_states: all_hidden_states += (hidden_states,) past_key_value = past_key_values[i] if past_key_values is not None else None if getattr(self, "gradient_checkpointing", False) and self.training: hidden_states, _ = checkpoint( mixer, hidden_states, attention_mask, position_ids, None, use_reentrant=False ) else: hidden_states, present = mixer(hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value) if use_cache: presents += (present,) hidden_states = self.ln_f(hidden_states) if output_hidden_states: all_hidden_states += (hidden_states,) logits = self.lm_head(hidden_states) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss_fct = nn.CrossEntropyLoss(ignore_index=-100) loss = loss_fct( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ) if not return_dict: output = (logits,) if output_hidden_states: output += (all_hidden_states,) return ((loss,) + output) if loss is not None else output return HybridTimeScaleOutput( loss=loss, logits=logits, past_key_values=presents, hidden_states=all_hidden_states, attentions=None, last_hidden_state=hidden_states, ) def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs): position_ids = kwargs.get("position_ids", None) if attention_mask is not None and position_ids is None: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 1) if past_key_values is not None: if isinstance(input_ids, torch.Tensor): input_ids = input_ids[:, -1:] if position_ids is not None: position_ids = position_ids[:, -1].unsqueeze(-1) model_inputs = { "input_ids": input_ids, "past_key_values": past_key_values, "use_cache": kwargs.get("use_cache", True), "position_ids": position_ids, "attention_mask": attention_mask, } return model_inputs def _prepare_cache_for_generation(self, *args, **kwargs): # Override GenerationMixin's method to bypass Hugging Face's DynamicCache initialization. # This completely avoids the KeyError: 'linear' crash by ensuring HF uses standard tuple caching. return None def _get_initial_cache(self, **kwargs): return None def _reorder_cache(self, past_key_values, beam_idx): return past_key_values # ── Register with AutoClasses ──────────────────────────────────────── AutoConfig.register("hybrid_timescale_lm", HybridTimeScaleConfig) AutoModelForCausalLM.register(HybridTimeScaleConfig, HybridTimeScaleLM) # Required for push_to_hub to upload the custom python code and generate auto_map HybridTimeScaleConfig.register_for_auto_class("AutoConfig") HybridTimeScaleLM.register_for_auto_class("AutoModelForCausalLM")