Download last-checkpoint/model.py from CodeIsAbstract/HybridTimeScale_1B: direct link, hf CLI and curl.
- Browser
- Download file 29.4 kB
-
https://huggingface.co/CodeIsAbstract/HybridTimeScale_1B/resolve/main/last-checkpoint/model.py
- Command line
-
hf download hf://CodeIsAbstract/HybridTimeScale_1B/last-checkpoint/model.py
-
curl -L -o model.py https://huggingface.co/CodeIsAbstract/HybridTimeScale_1B/resolve/main/last-checkpoint/model.py
29.4 kB
| """ | |
| 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 | |
| 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") | |