| """Hugging Face Transformers model for the Cubit KRR token mixer.""" |
|
|
| from __future__ import annotations |
|
|
| import math |
| from typing import Optional, Union |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from transformers import PreTrainedModel |
| from transformers.generation import GenerationMixin |
| from transformers.modeling_outputs import BaseModelOutput, CausalLMOutput |
|
|
| from .configuration_cubit import CubitConfig |
|
|
|
|
| def _batch_matmul(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor: |
| """Multiply ``[B, H, M, K]`` by ``[B, H, K, N]``.""" |
| batch_size, n_heads, _, inner_dim = left.shape |
| if right.shape[:2] != (batch_size, n_heads) or right.shape[2] != inner_dim: |
| raise RuntimeError("incompatible batched matrix multiplication shapes") |
| return torch.matmul(left, right) |
|
|
|
|
| def _score_mask( |
| *, |
| batch_size: int, |
| q_start: int, |
| q_end: int, |
| k_start: int, |
| k_end: int, |
| device: torch.device, |
| window_size: Optional[int], |
| ) -> torch.Tensor: |
| """Build one causal (and optionally sliding-window) mask tile.""" |
| query_positions = torch.arange(q_start, q_end, device=device) |
| key_positions = torch.arange(k_start, k_end, device=device) |
| mask = key_positions[None, :] <= query_positions[:, None] |
| if window_size is not None: |
| mask = mask & ( |
| key_positions[None, :] >= query_positions[:, None] - (window_size - 1) |
| ) |
| return mask.view(1, 1, q_end - q_start, k_end - k_start).expand( |
| batch_size, 1, -1, -1 |
| ) |
|
|
|
|
| def _kernel_weights( |
| reference_queries: torch.Tensor, |
| reference_keys: torch.Tensor, |
| *, |
| q_start: int, |
| k_start: int, |
| window_size: Optional[int], |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| """Return one causal Cubit kernel tile and its row log-normalizers.""" |
| q_end = q_start + reference_queries.shape[2] |
| k_end = k_start + reference_keys.shape[2] |
| scores = _batch_matmul(reference_queries, reference_keys.transpose(-2, -1)) |
| mask = _score_mask( |
| batch_size=reference_queries.shape[0], |
| q_start=q_start, |
| q_end=q_end, |
| k_start=k_start, |
| k_end=k_end, |
| device=reference_queries.device, |
| window_size=window_size, |
| ) |
| scores = scores.masked_fill(~mask, -torch.inf) |
| logsumexp = torch.logsumexp(scores, dim=-1) |
| weights = torch.exp(scores - logsumexp.unsqueeze(-1)).masked_fill(~mask, 0.0) |
| return weights, logsumexp |
|
|
|
|
| def _first_key_for_block(q_start: int, window_size: Optional[int]) -> int: |
| if window_size is None: |
| return 0 |
| return max(0, q_start - window_size + 1) |
|
|
|
|
| def streaming_causal_krr_solve( |
| reference_queries: torch.Tensor, |
| reference_keys: torch.Tensor, |
| rhs: torch.Tensor, |
| regularization: torch.Tensor, |
| *, |
| window_size: Optional[int], |
| block_size: int, |
| ) -> torch.Tensor: |
| """Run the exact block-streaming forward solve used by the OLMo Core Cubit model.""" |
| reference_queries = reference_queries.contiguous() |
| reference_keys = reference_keys.contiguous() |
| rhs = rhs.contiguous() |
| batch_size, n_heads, seq_len, _ = rhs.shape |
| solution = torch.empty_like(rhs) |
|
|
| for q_start in range(0, seq_len, block_size): |
| q_end = min(q_start + block_size, seq_len) |
| key_start = _first_key_for_block(q_start, window_size) |
| weights, _ = _kernel_weights( |
| reference_queries[:, :, q_start:q_end], |
| reference_keys[:, :, key_start:q_end], |
| q_start=q_start, |
| k_start=key_start, |
| window_size=window_size, |
| ) |
|
|
| previous_end = q_start - key_start |
| residual = rhs[:, :, q_start:q_end] |
| if previous_end > 0: |
| residual = residual - _batch_matmul( |
| weights[:, :, :, :previous_end], |
| solution[:, :, key_start:q_start], |
| ) |
|
|
| diagonal_block = weights[:, :, :, previous_end:] |
| block_len = q_end - q_start |
| identity = torch.eye( |
| block_len, |
| dtype=diagonal_block.dtype, |
| device=diagonal_block.device, |
| ).view(1, 1, block_len, block_len) |
| diagonal_block = diagonal_block + regularization.view(1, n_heads, 1, 1) * identity |
| solution[:, :, q_start:q_end] = torch.linalg.solve_triangular( |
| diagonal_block, |
| residual, |
| upper=False, |
| ) |
|
|
| return solution |
|
|
|
|
| class CubitRMSNorm(nn.Module): |
| """RMSNorm with the same fp32 calculation and cast order as OLMo Core.""" |
|
|
| def __init__(self, hidden_size: int, eps: float) -> 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.float() |
| variance = hidden_states.pow(2).mean(-1, keepdim=True) |
| hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) |
| |
| |
| |
| hidden_states = self.weight.type_as(hidden_states) * hidden_states |
| return hidden_states.to(input_dtype) |
|
|
|
|
| class CubitRotaryEmbedding(nn.Module): |
| """Full-precision RoPE matching OLMo Core's real-valued implementation.""" |
|
|
| def __init__(self, config: CubitConfig) -> None: |
| super().__init__() |
| self.dim = config.head_dim |
| self.theta = config.rope_theta |
|
|
| @staticmethod |
| def _rotate_half(hidden_states: torch.Tensor) -> torch.Tensor: |
| batch_size, seq_len, n_heads, head_dim = hidden_states.shape |
| hidden_states = hidden_states.view(batch_size, seq_len, n_heads, 2, head_dim // 2) |
| first, second = hidden_states.unbind(dim=-2) |
| return torch.cat((-second, first), dim=-1) |
|
|
| def _get_sin_cos( |
| self, |
| seq_len: int, |
| device: torch.device, |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| with torch.autocast(device.type, enabled=False): |
| inv_freq = 1.0 / ( |
| self.theta |
| ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float) / self.dim) |
| ) |
| positions = torch.arange(seq_len, device=device, dtype=torch.float) |
| frequencies = torch.einsum("i , j -> i j", positions, inv_freq) |
| embeddings = torch.cat((frequencies, frequencies), dim=-1) |
| return embeddings.sin(), embeddings.cos() |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| position_ids: torch.Tensor, |
| ) -> torch.Tensor: |
| input_dtype = hidden_states.dtype |
| hidden_states = hidden_states.float() |
| max_position = int(position_ids.max().item()) + 1 |
| pos_sin, pos_cos = self._get_sin_cos(max_position, hidden_states.device) |
| pos_sin = pos_sin[position_ids].unsqueeze(2).type_as(hidden_states) |
| pos_cos = pos_cos[position_ids].unsqueeze(2).type_as(hidden_states) |
| with torch.autocast(hidden_states.device.type, enabled=False): |
| hidden_states = (hidden_states * pos_cos) + ( |
| self._rotate_half(hidden_states) * pos_sin |
| ) |
| return hidden_states.to(input_dtype) |
|
|
|
|
| class CubitMLP(nn.Module): |
| """SwiGLU feed-forward layer.""" |
|
|
| def __init__(self, config: CubitConfig) -> None: |
| super().__init__() |
| self.gate_proj = nn.Linear( |
| config.hidden_size, config.intermediate_size, bias=config.mlp_bias |
| ) |
| self.up_proj = nn.Linear( |
| config.hidden_size, config.intermediate_size, bias=config.mlp_bias |
| ) |
| self.down_proj = nn.Linear( |
| config.intermediate_size, config.hidden_size, bias=config.mlp_bias |
| ) |
|
|
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: |
| return self.down_proj(F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states)) |
|
|
|
|
| class CubitAttention(nn.Module): |
| """Causal kernel-ridge-regression token mixer followed by output attention.""" |
|
|
| def __init__(self, config: CubitConfig, layer_idx: int) -> None: |
| super().__init__() |
| self.config = config |
| self.layer_idx = layer_idx |
| self.n_heads = config.num_attention_heads |
| self.head_dim = config.head_dim |
| self.hidden_size = config.hidden_size |
| self.window_size = ( |
| config.sliding_window if config.layer_types[layer_idx] == "sliding_attention" else None |
| ) |
| self.share_reference = config.share_reference |
|
|
| self.q_proj = nn.Linear( |
| config.hidden_size, config.hidden_size, bias=config.attention_bias |
| ) |
| self.k_proj = nn.Linear( |
| config.hidden_size, config.hidden_size, bias=config.attention_bias |
| ) |
| self.v_proj = nn.Linear( |
| config.hidden_size, config.hidden_size, bias=config.attention_bias |
| ) |
| self.o_proj = nn.Linear( |
| config.hidden_size, config.hidden_size, bias=config.attention_bias |
| ) |
| self.r_proj = ( |
| None |
| if config.share_reference |
| else nn.Linear(config.hidden_size, config.hidden_size, bias=config.attention_bias) |
| ) |
| self.lrr_proj = nn.Linear( |
| config.hidden_size, config.num_attention_heads, bias=config.attention_bias |
| ) |
| self.q_norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) |
| self.k_norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) |
| self.lrr_lower = nn.Parameter(torch.full((self.n_heads,), 0.5)) |
| self.lrr_range = nn.Parameter(torch.full((self.n_heads,), 1.5)) |
| self.reference_scale = nn.Parameter(torch.ones(self.n_heads)) |
| self.log_regularization = nn.Parameter( |
| torch.full((self.n_heads,), math.log(1e-10)) |
| ) |
| self.rotary_emb = CubitRotaryEmbedding(config) |
|
|
| def _build_attention_mask( |
| self, |
| batch_size: int, |
| seq_len: int, |
| device: torch.device, |
| attention_mask: Optional[torch.Tensor], |
| ) -> torch.Tensor: |
| positions = torch.arange(seq_len, device=device) |
| query_pos = positions[:, None] |
| key_pos = positions[None, :] |
| mask = key_pos <= query_pos |
| if self.window_size is not None: |
| mask = mask & (key_pos >= query_pos - (self.window_size - 1)) |
| mask = mask.unsqueeze(0).expand(batch_size, -1, -1) |
|
|
| if attention_mask is not None: |
| if attention_mask.ndim == 2: |
| key_mask = attention_mask.to(device=device, dtype=torch.bool) |
| mask = mask & key_mask[:, None, :] |
| elif attention_mask.ndim == 4: |
| layer_mask = attention_mask[:, 0] |
| allowed = layer_mask if layer_mask.dtype == torch.bool else layer_mask >= 0 |
| mask = mask & allowed.to(device=device) |
| else: |
| raise ValueError("attention_mask must have rank 2 or 4") |
|
|
| |
| empty_rows = ~mask.any(dim=-1) |
| if empty_rows.any(): |
| diagonal = torch.eye(seq_len, dtype=torch.bool, device=device).unsqueeze(0) |
| mask = mask | (empty_rows.unsqueeze(-1) & diagonal) |
| return mask |
|
|
| @staticmethod |
| def _masked_softmax(scores: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: |
| return torch.softmax(scores.masked_fill(~mask[:, None, :, :], -torch.inf), dim=-1) |
|
|
| def _dense_krr_solve( |
| self, |
| reference_queries: torch.Tensor, |
| reference_keys: torch.Tensor, |
| rhs: torch.Tensor, |
| regularization: torch.Tensor, |
| mask: torch.Tensor, |
| ) -> torch.Tensor: |
| seq_len = rhs.shape[2] |
| inverse_sigma = self._masked_softmax( |
| reference_queries @ reference_keys.transpose(-2, -1), mask |
| ) |
| identity = torch.eye(seq_len, device=rhs.device, dtype=torch.float32) |
| inverse_sigma = inverse_sigma + regularization.view( |
| 1, self.n_heads, 1, 1 |
| ) * identity.view(1, 1, seq_len, seq_len) |
| return torch.linalg.solve_triangular(inverse_sigma, rhs, upper=False) |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.Tensor] = None, |
| **kwargs, |
| ) -> torch.Tensor: |
| if kwargs.get("past_key_values") is not None or kwargs.get("past_key_value") is not None: |
| raise NotImplementedError("Cubit v1 does not support KV caching") |
| batch_size, seq_len, _ = hidden_states.shape |
| if position_ids is None: |
| position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0) |
| position_ids = position_ids.expand(batch_size, -1) |
|
|
| q = self.q_norm(self.q_proj(hidden_states)) |
| k = self.k_norm(self.k_proj(hidden_states)) |
| v = self.v_proj(hidden_states) |
| projected_r = None if self.share_reference else self.r_proj(hidden_states) |
|
|
| q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) |
| k = k.view(batch_size, seq_len, self.n_heads, self.head_dim) |
| v = v.view(batch_size, seq_len, self.n_heads, self.head_dim) |
| if self.share_reference: |
| r = k |
| else: |
| assert projected_r is not None |
| r = projected_r.view(batch_size, seq_len, self.n_heads, self.head_dim) |
|
|
| reference_scale = self.reference_scale.float().view(1, 1, self.n_heads, 1) |
| normalized_r = F.normalize( |
| r.float(), p=2, dim=-1, eps=self.config.reference_norm_eps |
| ) * reference_scale |
|
|
| q = self.rotary_emb(q, position_ids) |
| k = self.rotary_emb(k, position_ids) |
| r = self.rotary_emb(r, position_ids) |
| normalized_r = self.rotary_emb(normalized_r, position_ids) |
|
|
| r_heads = r.float().transpose(1, 2) |
| normalized_r_heads = normalized_r.float().transpose(1, 2) |
| lrr_logits = self.lrr_proj(hidden_states).float().transpose(1, 2).unsqueeze(-1) |
| lrr = self.lrr_lower.float().view(1, self.n_heads, 1, 1) |
| lrr = lrr + self.lrr_range.float().view(1, self.n_heads, 1, 1) * torch.sigmoid( |
| lrr_logits |
| ) |
| rhs = lrr * v.float().transpose(1, 2) |
| regularization = self.log_regularization.float().exp() |
|
|
| mask = self._build_attention_mask( |
| batch_size, seq_len, hidden_states.device, attention_mask |
| ) |
| use_streaming = self.config.krr_implementation == "streaming" |
| if attention_mask is not None and attention_mask.ndim == 2: |
| use_streaming = use_streaming and bool(attention_mask.to(torch.bool).all()) |
| if use_streaming: |
| solution = streaming_causal_krr_solve( |
| r_heads, |
| normalized_r_heads, |
| rhs, |
| regularization, |
| window_size=self.window_size, |
| block_size=self.config.krr_block_size, |
| ) |
| else: |
| solution = self._dense_krr_solve( |
| r_heads, normalized_r_heads, rhs, regularization, mask |
| ) |
| solution = solution.transpose(1, 2).to(q.dtype).contiguous() |
|
|
| scores = torch.einsum("bthd,bshd->bhts", q.float(), k.float()) * ( |
| self.head_dim**-0.5 |
| ) |
| weights = self._masked_softmax(scores, mask) |
| output = torch.einsum("bhts,bshd->bthd", weights, solution.float()).to(q.dtype) |
| return self.o_proj(output.reshape(batch_size, seq_len, -1)) |
|
|
|
|
| class CubitDecoderLayer(nn.Module): |
| """OLMo 3 reordered-norm decoder block.""" |
|
|
| def __init__(self, config: CubitConfig, layer_idx: int) -> None: |
| super().__init__() |
| self.self_attn = CubitAttention(config, layer_idx) |
| self.mlp = CubitMLP(config) |
| self.post_attention_layernorm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) |
| self.post_feedforward_layernorm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) |
|
|
| def forward( |
| self, |
| hidden_states: torch.Tensor, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.Tensor] = None, |
| **kwargs, |
| ) -> torch.Tensor: |
| hidden_states = hidden_states + self.post_attention_layernorm( |
| self.self_attn( |
| hidden_states, |
| attention_mask=attention_mask, |
| position_ids=position_ids, |
| **kwargs, |
| ) |
| ) |
| return hidden_states + self.post_feedforward_layernorm(self.mlp(hidden_states)) |
|
|
|
|
| class CubitPreTrainedModel(PreTrainedModel): |
| """Base class for Cubit Transformers models.""" |
|
|
| config_class = CubitConfig |
| base_model_prefix = "model" |
| supports_gradient_checkpointing = False |
| _no_split_modules = ["CubitDecoderLayer"] |
| _supports_flash_attn = False |
| _supports_sdpa = False |
| _supports_cache_class = False |
|
|
| def _init_weights(self, module: nn.Module) -> None: |
| if isinstance(module, nn.Linear): |
| module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) |
| if module.bias is not None: |
| module.bias.data.zero_() |
| elif isinstance(module, nn.Embedding): |
| module.weight.data.normal_(mean=0.0, std=self.config.initializer_range) |
| if module.padding_idx is not None: |
| module.weight.data[module.padding_idx].zero_() |
|
|
|
|
| class CubitModel(CubitPreTrainedModel): |
| """Bare Cubit decoder model.""" |
|
|
| def __init__(self, config: CubitConfig) -> None: |
| super().__init__(config) |
| self.padding_idx = config.pad_token_id |
| self.embed_tokens = nn.Embedding( |
| config.vocab_size, config.hidden_size, padding_idx=self.padding_idx |
| ) |
| self.layers = nn.ModuleList( |
| [CubitDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] |
| ) |
| self.norm = CubitRMSNorm(config.hidden_size, config.rms_norm_eps) |
| self.post_init() |
|
|
| def get_input_embeddings(self) -> nn.Embedding: |
| return self.embed_tokens |
|
|
| def set_input_embeddings(self, value: nn.Embedding) -> None: |
| self.embed_tokens = value |
|
|
| def forward( |
| self, |
| input_ids: Optional[torch.LongTensor] = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| inputs_embeds: Optional[torch.FloatTensor] = None, |
| use_cache: Optional[bool] = None, |
| output_attentions: Optional[bool] = None, |
| output_hidden_states: Optional[bool] = None, |
| return_dict: Optional[bool] = None, |
| **kwargs, |
| ) -> Union[tuple, BaseModelOutput]: |
| if (input_ids is None) == (inputs_embeds is None): |
| raise ValueError("provide exactly one of input_ids or inputs_embeds") |
| if use_cache: |
| raise NotImplementedError("Cubit v1 does not support KV caching") |
| if output_attentions: |
| raise NotImplementedError("Cubit v1 does not return attention matrices") |
| output_hidden_states = ( |
| self.config.output_hidden_states |
| if output_hidden_states is None |
| else output_hidden_states |
| ) |
| return_dict = self.config.use_return_dict if return_dict is None else return_dict |
|
|
| hidden_states = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds |
| batch_size, seq_len, _ = hidden_states.shape |
| if position_ids is None: |
| if attention_mask is not None and attention_mask.ndim == 2: |
| position_ids = attention_mask.long().cumsum(-1) - 1 |
| position_ids.masked_fill_(attention_mask == 0, 0) |
| else: |
| position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0) |
| position_ids = position_ids.expand(batch_size, -1) |
|
|
| collected_hidden_states = () if output_hidden_states else None |
| for decoder_layer in self.layers: |
| if collected_hidden_states is not None: |
| collected_hidden_states += (hidden_states,) |
| hidden_states = decoder_layer( |
| hidden_states, |
| attention_mask=attention_mask, |
| position_ids=position_ids, |
| **kwargs, |
| ) |
| hidden_states = self.norm(hidden_states) |
| if collected_hidden_states is not None: |
| collected_hidden_states += (hidden_states,) |
|
|
| if not return_dict: |
| return tuple(value for value in (hidden_states, collected_hidden_states) if value is not None) |
| return BaseModelOutput( |
| last_hidden_state=hidden_states, |
| hidden_states=collected_hidden_states, |
| attentions=None, |
| ) |
|
|
|
|
| class CubitForCausalLM(CubitPreTrainedModel, GenerationMixin): |
| """Cubit decoder with a causal language-modeling head.""" |
|
|
| _tied_weights_keys: list[str] = [] |
|
|
| def __init__(self, config: CubitConfig) -> None: |
| super().__init__(config) |
| self.model = CubitModel(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) -> nn.Embedding: |
| return self.model.embed_tokens |
|
|
| def set_input_embeddings(self, value: nn.Embedding) -> None: |
| self.model.embed_tokens = value |
|
|
| def get_output_embeddings(self) -> nn.Linear: |
| return self.lm_head |
|
|
| def set_output_embeddings(self, value: nn.Linear) -> None: |
| self.lm_head = value |
|
|
| def forward( |
| self, |
| input_ids: Optional[torch.LongTensor] = None, |
| attention_mask: Optional[torch.Tensor] = None, |
| position_ids: Optional[torch.LongTensor] = None, |
| inputs_embeds: Optional[torch.FloatTensor] = None, |
| labels: Optional[torch.LongTensor] = None, |
| use_cache: Optional[bool] = None, |
| output_attentions: Optional[bool] = None, |
| output_hidden_states: Optional[bool] = None, |
| return_dict: Optional[bool] = None, |
| logits_to_keep: Union[int, torch.Tensor] = 0, |
| **kwargs, |
| ) -> Union[tuple, CausalLMOutput]: |
| return_dict = self.config.use_return_dict if return_dict is None else return_dict |
| outputs = self.model( |
| input_ids=input_ids, |
| attention_mask=attention_mask, |
| position_ids=position_ids, |
| inputs_embeds=inputs_embeds, |
| use_cache=use_cache, |
| output_attentions=output_attentions, |
| output_hidden_states=output_hidden_states, |
| return_dict=True, |
| **kwargs, |
| ) |
| hidden_states = outputs.last_hidden_state |
| if isinstance(logits_to_keep, int): |
| if logits_to_keep: |
| hidden_states = hidden_states[:, -logits_to_keep:, :] |
| else: |
| hidden_states = hidden_states.gather( |
| 1, logits_to_keep.unsqueeze(-1).expand(-1, -1, hidden_states.size(-1)) |
| ) |
| logits = self.lm_head(hidden_states) |
|
|
| loss = None |
| if labels is not None: |
| shift_logits = logits[..., :-1, :].contiguous().float() |
| shift_labels = labels[..., 1:].contiguous().to(shift_logits.device) |
| loss = F.cross_entropy( |
| shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1) |
| ) |
| if not return_dict: |
| values = (logits, outputs.hidden_states, outputs.attentions) |
| return ((loss,) + values) if loss is not None else values |
| return CausalLMOutput( |
| loss=loss, |
| logits=logits, |
| hidden_states=outputs.hidden_states, |
| attentions=outputs.attentions, |
| ) |
|
|