# coding=utf-8 # Modeling code for the Kanana-2 PD-series (Qwen3 backbone with sliding/full # alternating attention and per-attention-type RoPE). # # Implementation strategy # ----------------------- # The architecture is identical to Qwen3 except that the rotary embedding # differs between full-attention and sliding-attention layers. We therefore: # * keep the exact Qwen3 layer/attention/MLP/RMSNorm code (copied here so the # module is self-contained for `trust_remote_code=True` loading), and # * instantiate two rotary embeddings — one per attention type — and dispatch # to the right one in each decoder layer based on `layer_types`. # # The trick for "two rotary embeddings driven by one shared config" follows the # Gemma3 pattern: deepcopy the config and overwrite `rope_theta` / `rope_scaling` # to whatever the corresponding `config.rope_parameters[attention_type]` says, # then construct a standard rotary embedding from it. import copy from typing import Callable, Optional, Union import torch from torch import nn from transformers.activations import ACT2FN from transformers.cache_utils import Cache, DynamicCache from transformers.generation import GenerationMixin from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask from transformers.modeling_flash_attention_utils import FlashAttentionKwargs from transformers.modeling_layers import GradientCheckpointingLayer from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, auto_docstring, can_return_tuple from transformers.utils.deprecation import deprecate_kwarg # ── Cross-version compatibility shims ────────────────────────────────────── # Feature-detection (not version-string compare) because the Kakao-patched # transformers 5.3.0 selectively backports newer APIs, so plain version # inequalities give wrong answers on patched builds. # # Three points of divergence we handle here: # # 1. ``transformers.utils.generic.check_model_inputs`` — added around stock # 5.5; absent on Kakao-patched 5.3. Fall back to a no-op decorator. # # 2. ``create_causal_mask`` / ``create_sliding_window_causal_mask`` kwargs: # - ``input_embeds`` accepted ≤5.5 (deprecation alias); removed ≥5.6 # - ``inputs_embeds`` accepted ≥5.3 (patched) / ≥5.5 (stock) # - ``cache_position`` accepted ≤5.8; removed ≥5.9 # We pick the right embeds-kwarg name and filter out any kwarg the # installed version doesn't take. # # 3. ``ROPE_INIT_FUNCTIONS`` registry: # - Stock ≥5.5 has ``'proportional'`` (renamed from ``'default'``) # - Kakao-patched 5.3 has neither ``'default'`` nor ``'proportional'`` # We supply a local fallback for the unscaled-RoPE init when the # registry is missing both keys. import inspect as _inspect_compat # noqa: E402 try: from transformers.utils.generic import check_model_inputs # noqa: F401 except ImportError: def check_model_inputs(fn): # type: ignore[no-redef] return fn _CAUSAL_MASK_PARAMS = set(_inspect_compat.signature(create_causal_mask).parameters) _MASK_EMBEDS_KW = ( "inputs_embeds" if "inputs_embeds" in _CAUSAL_MASK_PARAMS else "input_embeds" ) def _filter_mask_kwargs(kwargs: dict) -> dict: """Drop kwargs the installed ``create_causal_mask`` doesn't accept.""" return {k: v for k, v in kwargs.items() if k in _CAUSAL_MASK_PARAMS} def _compute_default_rope_inv_freq(config, device=None, seq_len=None): """Unscaled-RoPE inv_freq + attention scaling = 1.0. Mirrors transformers' canonical ``compute_default_rope_parameters`` — used when neither ``'default'`` nor ``'proportional'`` is in ``ROPE_INIT_FUNCTIONS``. """ if hasattr(config, "rope_parameters") and isinstance(config.rope_parameters, dict) \ and "rope_theta" in config.rope_parameters: base = config.rope_parameters["rope_theta"] else: base = getattr(config, "rope_theta", 10000.0) dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads inv_freq = 1.0 / ( base ** ( torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim ) ) return inv_freq, 1.0 def _resolve_rope_init(rope_type: str): """Pick a rope-init callable for ``rope_type`` across versions.""" if rope_type in ROPE_INIT_FUNCTIONS: return ROPE_INIT_FUNCTIONS[rope_type] # 'default' was renamed 'proportional' in stock ≥5.5 — try the other name. if rope_type == "default" and "proportional" in ROPE_INIT_FUNCTIONS: return ROPE_INIT_FUNCTIONS["proportional"] if rope_type == "proportional" and "default" in ROPE_INIT_FUNCTIONS: return ROPE_INIT_FUNCTIONS["default"] if rope_type in ("default", "proportional"): return _compute_default_rope_inv_freq raise KeyError( f"rope_type={rope_type!r} not in ROPE_INIT_FUNCTIONS and no fallback " f"available; keys={sorted(ROPE_INIT_FUNCTIONS)}" ) del _inspect_compat # ─────────────────────────────────────────────────────────────────────────── from .configuration_kanana2_tiny import Kanana2TinyConfig # --------------------------------------------------------------------------- # Building blocks (copied verbatim from Qwen3) # --------------------------------------------------------------------------- class Kanana2TinyRMSNorm(nn.Module): def __init__(self, hidden_size, eps: float = 1e-6) -> None: super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) return self.weight * hidden_states.to(input_dtype) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" class Kanana2TinyMLP(nn.Module): def __init__(self, config): super().__init__() self.config = config self.hidden_size = config.hidden_size self.intermediate_size = config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.act_fn = ACT2FN[config.hidden_act] def forward(self, x): return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) def rotate_half(x): x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: batch, num_key_value_heads, slen, head_dim = hidden_states.shape if n_rep == 1: return hidden_states hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) def eager_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], scaling: float, dropout: float = 0.0, **kwargs: Unpack[TransformersKwargs], ): key_states = repeat_kv(key, module.num_key_value_groups) value_states = repeat_kv(value, module.num_key_value_groups) attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling if attention_mask is not None: causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] attn_weights = attn_weights + causal_mask attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) attn_output = torch.matmul(attn_weights, value_states) attn_output = attn_output.transpose(1, 2).contiguous() return attn_output, attn_weights class Kanana2TinyAttention(nn.Module): """Multi-headed attention (identical to Qwen3Attention).""" def __init__(self, config: Kanana2TinyConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads) self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads self.scaling = self.head_dim**-0.5 self.attention_dropout = config.attention_dropout self.is_causal = True self.q_proj = nn.Linear( config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias ) self.k_proj = nn.Linear( config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias ) self.v_proj = nn.Linear( config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias ) self.o_proj = nn.Linear( config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias ) self.q_norm = Kanana2TinyRMSNorm(self.head_dim, eps=config.rms_norm_eps) self.k_norm = Kanana2TinyRMSNorm(self.head_dim, eps=config.rms_norm_eps) self.sliding_window = config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58") def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor], past_key_values: Optional[Cache] = None, cache_position: Optional[torch.LongTensor] = None, **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, Optional[torch.Tensor]]: input_shape = hidden_states.shape[:-1] hidden_shape = (*input_shape, -1, self.head_dim) query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2) key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2) value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2) cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) if past_key_values is not None: cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) attention_interface: Callable = eager_attention_forward if self.config._attn_implementation != "eager": attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] attn_output, attn_weights = attention_interface( self, query_states, key_states, value_states, attention_mask, dropout=0.0 if not self.training else self.attention_dropout, scaling=self.scaling, sliding_window=self.sliding_window, **kwargs, ) attn_output = attn_output.reshape(*input_shape, -1).contiguous() attn_output = self.o_proj(attn_output) return attn_output, attn_weights class Kanana2TinyDecoderLayer(GradientCheckpointingLayer): def __init__(self, config: Kanana2TinyConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size self.self_attn = Kanana2TinyAttention(config=config, layer_idx=layer_idx) self.mlp = Kanana2TinyMLP(config) self.input_layernorm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.post_attention_layernorm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.attention_type = config.layer_types[layer_idx] @deprecate_kwarg("past_key_value", new_name="past_key_values", version="4.58") def forward( self, hidden_states: torch.Tensor, position_embeddings_full: tuple[torch.Tensor, torch.Tensor], position_embeddings_sliding: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[Cache] = None, use_cache: Optional[bool] = False, cache_position: Optional[torch.LongTensor] = None, **kwargs: Unpack[TransformersKwargs], ) -> torch.Tensor: # Pick the right RoPE for this layer type. if self.attention_type == "sliding_attention": position_embeddings = position_embeddings_sliding else: position_embeddings = position_embeddings_full residual = hidden_states hidden_states = self.input_layernorm(hidden_states) hidden_states, _ = self.self_attn( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, **kwargs, ) hidden_states = residual + hidden_states residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) hidden_states = residual + hidden_states return hidden_states # --------------------------------------------------------------------------- # Rotary embedding (driven by `config.rope_scaling` for the chosen attn type) # --------------------------------------------------------------------------- class Kanana2TinyRotaryEmbedding(nn.Module): """Standard Qwen3-style rotary embedding. The per-attention-type difference is encoded by the *config* passed in: callers construct two of these from views built via `_make_attention_specific_config` below. """ inv_freq: torch.Tensor def __init__(self, config: Kanana2TinyConfig, device=None): super().__init__() # BC: "rope_type" was originally "type" if hasattr(config, "rope_scaling") and isinstance(config.rope_scaling, dict): self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type", "default")) else: self.rope_type = "default" self.max_seq_len_cached = config.max_position_embeddings self.original_max_seq_len = config.max_position_embeddings self.config = config # Resolve rope init across stock 5.4 (had 'default'), stock 5.5+ # (renamed to 'proportional'), and Kakao-patched 5.3 (has neither; # falls through to our local unscaled-RoPE impl). self.rope_init_fn = _resolve_rope_init(self.rope_type) inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device) self.register_buffer("inv_freq", inv_freq, persistent=False) self.original_inv_freq = self.inv_freq @staticmethod def compute_default_rope_parameters(config, device=None, seq_len=None): """Stock transformers ≥5.9's ``modeling_utils._init_weights`` calls ``module.compute_default_rope_parameters`` directly when ``rope_type == "default"`` (instead of looking it up in ``ROPE_INIT_FUNCTIONS``). This staticmethod has to exist on the class for that init pass to find it; the body is the same unscaled inv_freq computation we use as a fallback elsewhere. """ return _compute_default_rope_inv_freq(config, device=device, seq_len=seq_len) @torch.no_grad() @dynamic_rope_update def forward(self, x, position_ids): inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device) position_ids_expanded = position_ids[:, None, :].float() device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu" with torch.autocast(device_type=device_type, enabled=False): freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() * self.attention_scaling sin = emb.sin() * self.attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) def _make_attention_specific_config(config: Kanana2TinyConfig, attention_type: str): """Return a deep copy of `config` configured for a single attention type's RoPE. 4.57.1's `ROPE_INIT_FUNCTIONS` read `config.rope_theta` (top-level) and `config.rope_scaling` (a flat dict with `rope_type`/`factor`/...), so we flatten `config.rope_parameters[attention_type]` into that shape: pop `rope_theta` up to the top level, and leave the remaining keys in `rope_scaling`. For `rope_type='default'` this leaves a 1-key `{"rope_type": "default"}` dict, which `_validate_default_rope_parameters` accepts cleanly. """ if attention_type not in config.rope_parameters: raise KeyError( f"rope_parameters is missing entry for attention_type={attention_type!r}; " f"available keys: {list(config.rope_parameters.keys())}" ) params = dict(config.rope_parameters[attention_type]) new_config = copy.deepcopy(config) new_config.rope_theta = params.pop("rope_theta", config.rope_theta) new_config.rope_scaling = params return new_config # --------------------------------------------------------------------------- # Pretrained model classes # --------------------------------------------------------------------------- @auto_docstring class Kanana2TinyPreTrainedModel(PreTrainedModel): config: Kanana2TinyConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["Kanana2TinyDecoderLayer"] _skip_keys_device_placement = ["past_key_values"] _supports_flash_attn = True _supports_sdpa = True _supports_flex_attn = True _can_compile_fullgraph = True _supports_attention_backend = True _can_record_outputs = { "hidden_states": Kanana2TinyDecoderLayer, "attentions": Kanana2TinyAttention, } @auto_docstring class Kanana2TinyModel(Kanana2TinyPreTrainedModel): def __init__(self, config: Kanana2TinyConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList( [Kanana2TinyDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self.norm = Kanana2TinyRMSNorm(config.hidden_size, eps=config.rms_norm_eps) # Two rotary embeddings, one per attention type. See the Gemma3 # implementation for the same pattern. full_cfg = _make_attention_specific_config(config, "full_attention") self.rotary_emb_full = Kanana2TinyRotaryEmbedding(config=full_cfg) if "sliding_attention" in config.layer_types: sliding_cfg = _make_attention_specific_config(config, "sliding_attention") self.rotary_emb_sliding = Kanana2TinyRotaryEmbedding(config=sliding_cfg) else: self.rotary_emb_sliding = None # Backward-compat alias so any helper that expects `model.rotary_emb` # (e.g. some training-time monkey patches) still finds something. self.rotary_emb = self.rotary_emb_full self.gradient_checkpointing = False self.has_sliding_layers = "sliding_attention" in config.layer_types self.post_init() @check_model_inputs @auto_docstring 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[Cache] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, **kwargs: Unpack[TransformersKwargs], ) -> BaseModelOutputWithPast: r""" cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*): Indices depicting the position of the input sequence tokens in the sequence. Used to update the cache in the correct position and to infer the complete sequence length. """ if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("You must specify exactly one of input_ids or inputs_embeds") if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) if use_cache and past_key_values is None: past_key_values = DynamicCache(config=self.config) if cache_position is None: past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 cache_position = torch.arange( past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device ) if position_ids is None: position_ids = cache_position.unsqueeze(0) if not isinstance(causal_mask_mapping := attention_mask, dict): mask_kwargs = _filter_mask_kwargs({ "config": self.config, _MASK_EMBEDS_KW: inputs_embeds, "attention_mask": attention_mask, "cache_position": cache_position, "past_key_values": past_key_values, "position_ids": position_ids, }) causal_mask_mapping = { "full_attention": create_causal_mask(**mask_kwargs), } if self.has_sliding_layers: causal_mask_mapping["sliding_attention"] = create_sliding_window_causal_mask(**mask_kwargs) hidden_states = inputs_embeds position_embeddings_full = self.rotary_emb_full(hidden_states, position_ids) if self.rotary_emb_sliding is not None: position_embeddings_sliding = self.rotary_emb_sliding(hidden_states, position_ids) else: position_embeddings_sliding = position_embeddings_full for decoder_layer in self.layers[: self.config.num_hidden_layers]: hidden_states = decoder_layer( hidden_states, position_embeddings_full=position_embeddings_full, position_embeddings_sliding=position_embeddings_sliding, attention_mask=causal_mask_mapping[decoder_layer.attention_type], position_ids=position_ids, past_key_values=past_key_values, use_cache=use_cache, cache_position=cache_position, **kwargs, ) hidden_states = self.norm(hidden_states) return BaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=past_key_values if use_cache else None, ) @auto_docstring class Kanana2TinyForCausalLM(Kanana2TinyPreTrainedModel, GenerationMixin): # transformers v5 changed this from list to dict (mapping tied-key -> source-key). # The list form still works on v4. Use the dict form for forward-compatibility. _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} _tp_plan = {"lm_head": "colwise_rep"} _pp_plan = {"lm_head": (["hidden_states"], ["logits"])} def __init__(self, config): super().__init__(config) self.model = Kanana2TinyModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.post_init() @can_return_tuple @auto_docstring 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[Cache] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, logits_to_keep: Union[int, torch.Tensor] = 0, **kwargs: Unpack[TransformersKwargs], ) -> CausalLMOutputWithPast: r""" cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*): Indices depicting the position of the input sequence tokens in the sequence. Used to update the cache in the correct position and to infer the complete sequence length. labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored (masked); the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`. """ outputs: BaseModelOutputWithPast = 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, cache_position=cache_position, **kwargs, ) hidden_states = outputs.last_hidden_state slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep logits = self.lm_head(hidden_states[:, slice_indices, :]) loss = None if labels is not None: loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, # Kanana2TinyModel.forward doesn't accumulate per-layer hidden_states even when # output_hidden_states=True; fall back to a 1-tuple of last_hidden_state so consumers # that index `hidden_states[-1]` (e.g. trl AutoModelForCausalLMWithValueHead) don't crash. hidden_states=outputs.hidden_states if outputs.hidden_states is not None else (outputs.last_hidden_state,), attentions=outputs.attentions, ) __all__ = [ "Kanana2TinyConfig", "Kanana2TinyForCausalLM", "Kanana2TinyModel", "Kanana2TinyPreTrainedModel", ]