CodeIsAbstract's picture
Training in progress, step 100, checkpoint
66f7fef verified
Raw History Blame Contribute Delete
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
@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")