Files
kanana-2-1.3b-instruct/modeling_kanana2_tiny.py

626 lines
27 KiB
Python
Raw Normal View History

# 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",
]