# coding=utf-8 # Copyright 2025 The Wiola / OSCOWL-AI authors. Apache-2.0. # # IMPORTANT: This model uses a custom 4‑tuple past_key_values # (k, v, cumsum, count). It is **incompatible** with the new # `DynamicCache` introduced in transformers ≥ 4.47. # Please pin your environment to `transformers==4.46.3`. """PyTorch Wiola model.""" from typing import List, Optional, Tuple, Union import torch import torch.nn as nn from transformers.generation import GenerationMixin from transformers.modeling_outputs import ( BaseModelOutputWithPast, CausalLMOutputWithPast, ) from transformers.modeling_utils import PreTrainedModel from .components.atm import merge_ratio, merge_tokens, unmerge_tokens from .components.dsff import DualStreamFeedForward from .components.gcla import GatedCrossLayerAttention from .components.normalization import WiolaRMSNorm from .configuration_wiola import WiolaConfig def _build_additive_mask(q_len, kv_len, device, dtype, key_padding=None): """Causal additive attention mask of shape [1, 1, q_len, kv_len]. key_padding: optional [B, kv_len] with 1 = keep, 0 = pad. Returns [B, 1, q_len, kv_len] if key_padding given, else [1,1,q_len,kv_len]. """ min_val = torch.finfo(dtype).min i = torch.arange(q_len, device=device)[:, None] j = torch.arange(kv_len, device=device)[None, :] allowed = j <= (kv_len - q_len + i) mask = torch.where( allowed, torch.zeros((), dtype=dtype, device=device), torch.full((), min_val, dtype=dtype, device=device), ) mask = mask[None, None] # [1,1,q,kv] if key_padding is not None: pad = (1 - key_padding[:, None, None, :].to(dtype)) * min_val mask = mask + pad return mask class WiolaDecoderLayer(nn.Module): def __init__(self, config: WiolaConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx self.input_norm = WiolaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.attn = GatedCrossLayerAttention(config, layer_idx) self.post_attn_norm = WiolaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.ffn = DualStreamFeedForward( config.hidden_size, config.dsff_narrow_size, config.dsff_wide_size ) # ATM is active during training in the middle third of the stack. lo = config.num_hidden_layers // 3 hi = 2 * config.num_hidden_layers // 3 self.atm_layer = lo <= layer_idx < hi self.last_merge_ratio = 0.0 def _run_attention( self, hidden_states, position_ids, attn_mask, context_summaries, past_key_value, use_cache ): return self.attn( hidden_states=hidden_states, position_ids=position_ids, attention_mask=attn_mask, context_summaries=context_summaries, past_key_value=past_key_value, use_cache=use_cache, ) def forward( self, hidden_states, position_ids, attn_mask, context_summaries=None, past_key_value=None, use_cache=False, ): residual = hidden_states normed = self.input_norm(hidden_states) atm_active = ( self.training and self.config.atm_enabled and self.atm_layer and past_key_value is None and normed.shape[1] >= 2 ) if atm_active: merged, keep_mask, merge_maps = merge_tokens(normed, self.config.atm_threshold) self.last_merge_ratio = merge_ratio(merge_maps, normed.shape[1]) bsz, t_prime, _ = merged.shape # Context gathered at each merged token's last source position. ctx_merged = None if context_summaries is not None and context_summaries.shape[2] > 0: last_idx = torch.zeros(bsz, t_prime, dtype=torch.long, device=merged.device) for b, groups in enumerate(merge_maps): for k, grp in enumerate(groups): last_idx[b, k] = grp[-1] batch_ar = torch.arange(bsz, device=merged.device)[:, None] ctx_merged = context_summaries[batch_ar, last_idx] # [B,T',Lam,d] m_mask = _build_additive_mask( t_prime, t_prime, merged.device, merged.dtype, key_padding=keep_mask ) m_pos = torch.arange(t_prime, device=merged.device)[None].expand(bsz, -1) attn_out_m, _ = self._run_attention(merged, m_pos, m_mask, ctx_merged, None, False) attn_out = unmerge_tokens(attn_out_m, merge_maps, normed.shape[1]) present = None else: self.last_merge_ratio = 0.0 attn_out, present = self._run_attention( normed, position_ids, attn_mask, context_summaries, past_key_value, use_cache ) hidden_states = residual + attn_out # Feed-forward block. residual = hidden_states normed = self.post_attn_norm(hidden_states) hidden_states = residual + self.ffn(normed) return hidden_states, present class WiolaPreTrainedModel(PreTrainedModel): config_class = WiolaConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["WiolaDecoderLayer"] _skip_keys_device_placement = "past_key_values" def _init_weights(self, module): std = self.config.initializer_range if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=std) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=std) if module.padding_idx is not None: module.weight.data[module.padding_idx].zero_() elif isinstance(module, WiolaRMSNorm): module.weight.data.fill_(1.0) module.offset.data.zero_() class WiolaModel(WiolaPreTrainedModel): def __init__(self, config: WiolaConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList( [WiolaDecoderLayer(config, i) for i in range(config.num_hidden_layers)] ) self.final_norm = WiolaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.gradient_checkpointing = False self.lookback = config.gcla_lookback self.post_init() def get_input_embeddings(self): return self.embed_tokens def set_input_embeddings(self, value): self.embed_tokens = value @staticmethod def _layer_cummean(layer_out, past_sum, past_count): """Causal cumulative mean of layer_out over the sequence dim. layer_out: [B, S, d]; past_sum: [B, d] or None; past_count: [B,1] or None. Returns (cummean [B,S,d], new_sum [B,d], new_count [B,1]). """ bsz, s_len, dim = layer_out.shape if past_sum is None: past_sum = layer_out.new_zeros(bsz, dim) past_count = layer_out.new_zeros(bsz, 1) csum = past_sum[:, None, :] + torch.cumsum(layer_out, dim=1) # [B,S,d] steps = torch.arange(1, s_len + 1, device=layer_out.device, dtype=layer_out.dtype) counts = past_count[:, :, None] + steps[None, :, None] # [B,S,1] cummean = csum / counts.clamp_min(1.0) new_sum = past_sum + layer_out.sum(dim=1) # fix: use tensor creation to keep device/dtype consistent new_count = past_count + layer_out.new_tensor(float(s_len)) return cummean, new_sum, new_count def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[Tuple]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ): use_cache = use_cache if use_cache is not None else self.config.use_cache return_dict = return_dict if return_dict is not None else self.config.use_return_dict if input_ids is not None and inputs_embeds is not None: raise ValueError("Specify exactly one of input_ids or inputs_embeds.") if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) bsz, seq_len, _ = inputs_embeds.shape past_len = 0 # Only access past_key_values[0] if we can safely do so. if ( past_key_values is not None and len(past_key_values) > 0 and past_key_values[0] is not None and isinstance(past_key_values[0], tuple) and len(past_key_values[0]) >= 2 # at least (k,v) present and past_key_values[0][0] is not None ): past_len = past_key_values[0][0].shape[2] if position_ids is None: position_ids = torch.arange(past_len, past_len + seq_len, device=inputs_embeds.device)[ None ].expand(bsz, -1) kv_len = past_len + seq_len attn_mask = _build_additive_mask( seq_len, kv_len, inputs_embeds.device, inputs_embeds.dtype, key_padding=attention_mask, ) if self.gradient_checkpointing and self.training and use_cache: use_cache = False hidden_states = inputs_embeds prefix_means: List[torch.Tensor] = [] # cummean of each layer output next_cache: List[Tuple] = [] if use_cache else None for idx, layer in enumerate(self.layers): # Build per-position context from the most recent <= Lambda layers. ctx = None if prefix_means: take = prefix_means[-self.lookback :] ctx = torch.stack(take, dim=2) # [B, S, lam, d] past_kv = None past_sum = past_count = None # fix: guard against indexing past_key_values out of range if ( past_key_values is not None and idx < len(past_key_values) and past_key_values[idx] is not None ): pk = past_key_values[idx] # pk is expected to be a 4‑tuple (k, v, cumsum, count) if len(pk) == 4: past_kv = (pk[0], pk[1]) past_sum, past_count = pk[2], pk[3] else: # fallback for plain (k,v) cache – cannot recover cumsum, # so we start fresh (this will break recurrence but won't crash). past_kv = (pk[0], pk[1]) if self.gradient_checkpointing and self.training: hidden_states, present = self._gc_layer( layer, hidden_states, position_ids, attn_mask, ctx, past_kv, use_cache ) else: hidden_states, present = layer( hidden_states, position_ids, attn_mask, ctx, past_kv, use_cache ) cummean, new_sum, new_count = self._layer_cummean(hidden_states, past_sum, past_count) prefix_means.append(cummean) if use_cache: if present is None: next_cache.append(None) else: k, v = present next_cache.append((k, v, new_sum, new_count)) hidden_states = self.final_norm(hidden_states) if not return_dict: return (hidden_states, next_cache) return BaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=next_cache, ) def _gc_layer(self, layer, hidden_states, position_ids, attn_mask, ctx, past_kv, use_cache): def custom(hs): return layer(hs, position_ids, attn_mask, ctx, past_kv, use_cache) return torch.utils.checkpoint.checkpoint(custom, hidden_states, use_reentrant=False) class WiolaForCausalLM(WiolaPreTrainedModel, GenerationMixin): _tied_weights_keys = ["lm_head.weight"] def __init__(self, config: WiolaConfig): super().__init__(config) self.model = WiolaModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.post_init() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new): self.lm_head = new def get_decoder(self): return self.model def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[Tuple]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: return_dict = return_dict if return_dict is not None else self.config.use_return_dict outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, use_cache=use_cache, return_dict=True, ) hidden_states = outputs.last_hidden_state logits = self.lm_head(hidden_states).float() loss = None if labels is not None: shift_logits = logits[:, :-1, :].contiguous() shift_labels = labels[:, 1:].contiguous() loss = nn.functional.cross_entropy( shift_logits.view(-1, self.vocab_size), shift_labels.view(-1), ignore_index=-100, ) if not return_dict: out = (logits,) + (outputs.past_key_values,) return ((loss,) + out) if loss is not None else out return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, ) # --- Generation plumbing for the custom tuple cache -------------------- def prepare_inputs_for_generation( self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, **kwargs, ): has_past = ( past_key_values is not None and len(past_key_values) > 0 and past_key_values[0] is not None and isinstance(past_key_values[0], tuple) and len(past_key_values[0]) >= 2 and past_key_values[0][0] is not None ) if has_past: input_ids = input_ids[:, -1:] position_ids = kwargs.get("position_ids") if position_ids is None and attention_mask is not None: position_ids = attention_mask.long().cumsum(-1) - 1 position_ids.masked_fill_(attention_mask == 0, 1) if has_past: position_ids = position_ids[:, -input_ids.shape[1] :] return { "input_ids": input_ids, "past_key_values": past_key_values, "use_cache": kwargs.get("use_cache", True), "attention_mask": attention_mask, "position_ids": position_ids, } @staticmethod def _reorder_cache(past_key_values, beam_idx): if past_key_values is None: return None reordered = [] for layer in past_key_values: # fix: handle layers that are None (e.g. from ATM) if layer is None: reordered.append(None) continue k, v, s, c = layer reordered.append( ( k.index_select(0, beam_idx.to(k.device)), v.index_select(0, beam_idx.to(v.device)), s.index_select(0, beam_idx.to(s.device)), c.index_select(0, beam_idx.to(c.device)), ) ) return tuple(reordered)