scatterbrain-sm-experimental / modeling_scatterbrain.py
ToastyPigeon's picture
Upload folder using huggingface_hub
9711198 verified
Raw
History Blame Contribute Delete
28.5 kB
# Scatterbrain - A single MoE layer looped multiple times
#
# Based on Qwen3MoeForCausalLM architecture with looping logic from Loopstral
import copy
from typing import Callable, Optional, Union
import torch
import torch.nn.functional as F
from torch import nn
from transformers.activations import ACT2FN
from transformers.cache_utils import Cache, DynamicCache, DynamicLayer
# Try to import Liger kernels for optimized operations
try:
from liger_kernel.transformers import LigerRMSNorm, LigerSwiGLUMLP
LIGER_AVAILABLE = True
except ImportError:
LIGER_AVAILABLE = False
# Try to import ScatterMoE for optimized MoE computation
try:
from scattermoe import flatten_sort_count, parallel_linear
SCATTERMOE_AVAILABLE = True
except ImportError:
SCATTERMOE_AVAILABLE = False
from transformers.generation import GenerationMixin
from transformers.integrations import use_kernel_forward_from_hub
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 MoeCausalLMOutputWithPast, MoeModelOutputWithPast
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
from .configuration_scatterbrain import ScatterbrainConfig
def rotate_half(x):
"""Rotates half the hidden dims of the input."""
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):
"""Applies Rotary Position Embedding to the query and key tensors."""
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:
"""Repeat KV heads for GQA."""
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 ScatterbrainAttention(nn.Module):
"""Multi-headed attention with QK normalization."""
def __init__(self, config: ScatterbrainConfig, 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 = ScatterbrainRMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = ScatterbrainRMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.sliding_window = getattr(config, "sliding_window", 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,
cache_slot_idx: Optional[int] = 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}
# Use cache_slot_idx for looped layers - each iteration gets its own cache slot
slot_idx = cache_slot_idx if cache_slot_idx is not None else self.layer_idx
key_states, value_states = past_key_values.update(key_states, value_states, slot_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 ScatterbrainMLP(nn.Module):
def __init__(self, config, intermediate_size=None):
super().__init__()
self.config = config
self.hidden_size = config.hidden_size
self.intermediate_size = intermediate_size if intermediate_size is not None else 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):
# Fused SwiGLU: down(act(gate(x)) * up(x))
return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
class ScatterbrainSparseMoeBlock(nn.Module):
"""
MoE block with optimized expert computation using ScatterMoE Triton kernels.
When ScatterMoE is available, uses fused Triton kernels for massive speedup.
Falls back to Python loop implementation when ScatterMoE is not available.
Key optimizations with ScatterMoE:
1. Fused scatter-gather operations in Triton
2. All experts processed in parallel on GPU
3. Efficient memory access patterns
"""
def __init__(self, config):
super().__init__()
self.num_experts = config.num_experts
self.top_k = config.num_experts_per_tok
self.norm_topk_prob = config.norm_topk_prob
self.hidden_size = config.hidden_size
self.moe_intermediate_size = config.moe_intermediate_size
self.gate = nn.Linear(config.hidden_size, config.num_experts, bias=False)
# Stacked expert weights as Parameters
# Shape: [num_experts, out_features, in_features]
self.expert_gate_proj = nn.Parameter(
torch.empty(config.num_experts, config.moe_intermediate_size, config.hidden_size)
)
self.expert_up_proj = nn.Parameter(
torch.empty(config.num_experts, config.moe_intermediate_size, config.hidden_size)
)
self.expert_down_proj = nn.Parameter(
torch.empty(config.num_experts, config.hidden_size, config.moe_intermediate_size)
)
# Initialize weights
for param in [self.expert_gate_proj, self.expert_up_proj, self.expert_down_proj]:
nn.init.kaiming_uniform_(param, a=5**0.5)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, sequence_length, hidden_dim = hidden_states.shape
hidden_states_flat = hidden_states.view(-1, hidden_dim)
num_tokens = hidden_states_flat.shape[0]
# Router
router_logits = self.gate(hidden_states_flat)
routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)
routing_weights, selected_experts = torch.topk(routing_weights, self.top_k, dim=-1)
if self.norm_topk_prob:
routing_weights = routing_weights / routing_weights.sum(dim=-1, keepdim=True)
routing_weights = routing_weights.to(hidden_states_flat.dtype)
if SCATTERMOE_AVAILABLE:
# Use ScatterMoE's optimized Triton kernels
final_hidden_states = self._forward_scattermoe(
hidden_states_flat, selected_experts, routing_weights
)
else:
# Fallback to Python loop
final_hidden_states = self._forward_loop(
hidden_states_flat, selected_experts, routing_weights
)
final_hidden_states = final_hidden_states.reshape(batch_size, sequence_length, hidden_dim)
return final_hidden_states, router_logits
def _forward_scattermoe(
self,
hidden_states: torch.Tensor,
selected_experts: torch.Tensor,
routing_weights: torch.Tensor,
) -> torch.Tensor:
"""Forward pass using ScatterMoE Triton kernels."""
# Get sorting indices and expert offsets from ScatterMoE
sorted_expert_idxs, sorted_scattered_idxs, expert_offsets = flatten_sort_count(
selected_experts, num_experts=self.num_experts
)
# ScatterMoE parallel_linear expects weights in [E, in, out] format
# Our weights are stored as [E, out, in], so we need to permute them
# This matches what ParallelExperts.forward() does internally
# First pass: gate and up projections (hidden -> intermediate)
# We need to do gate and up separately, then combine with SiLU
gate_out = parallel_linear(
inputs=hidden_states,
expert_weights=self.expert_gate_proj.permute(0, 2, 1), # [E, hidden, intermediate]
k=self.top_k,
sorted_expert_idxs=sorted_expert_idxs,
sorted_scattered_idxs=sorted_scattered_idxs,
expert_offsets=expert_offsets,
grouped_out=True, # Keep output grouped by expert (size: num_tokens * k)
)
up_out = parallel_linear(
inputs=hidden_states,
expert_weights=self.expert_up_proj.permute(0, 2, 1), # [E, hidden, intermediate]
k=self.top_k,
sorted_expert_idxs=sorted_expert_idxs,
sorted_scattered_idxs=sorted_scattered_idxs,
expert_offsets=expert_offsets,
grouped_out=True, # Keep output grouped by expert (size: num_tokens * k)
)
# Apply SwiGLU activation
activated = F.silu(gate_out) * up_out
# Second pass: down projection (intermediate -> hidden) with gates for weighted sum
# When grouped_in=True, use k=1 since input is already expanded to num_tokens * top_k
output = parallel_linear(
inputs=activated,
expert_weights=self.expert_down_proj.permute(0, 2, 1), # [E, intermediate, hidden]
k=1, # k=1 because input is already grouped/expanded
sorted_expert_idxs=sorted_expert_idxs,
sorted_scattered_idxs=sorted_scattered_idxs,
expert_offsets=expert_offsets,
gates=routing_weights, # Apply routing weights during scatter
grouped_in=True, # Input is grouped by expert
grouped_out=False, # Output ungrouped (back to token order)
)
return output
def _forward_loop(
self,
hidden_states: torch.Tensor,
selected_experts: torch.Tensor,
routing_weights: torch.Tensor,
) -> torch.Tensor:
"""Fallback forward pass using Python loop."""
num_tokens = hidden_states.shape[0]
# Flatten token-expert assignments
flat_expert_indices = selected_experts.view(-1)
flat_token_indices = torch.arange(num_tokens, device=hidden_states.device).unsqueeze(1).expand(-1, self.top_k).reshape(-1)
flat_routing_weights = routing_weights.view(-1)
# Sort by expert for contiguous memory access
sorted_indices = torch.argsort(flat_expert_indices, stable=True)
sorted_expert_indices = flat_expert_indices[sorted_indices]
sorted_token_indices = flat_token_indices[sorted_indices]
sorted_routing_weights = flat_routing_weights[sorted_indices]
sorted_hidden = hidden_states[sorted_token_indices]
# Compute expert boundaries
expert_counts = torch.bincount(sorted_expert_indices, minlength=self.num_experts)
expert_offsets = torch.zeros(self.num_experts + 1, dtype=torch.long, device=hidden_states.device)
expert_offsets[1:] = torch.cumsum(expert_counts, dim=0)
final_hidden_states = torch.zeros_like(hidden_states)
# Pre-fetch to lists to avoid repeated tensor->scalar conversions
counts_list = expert_counts.tolist()
offsets_list = expert_offsets.tolist()
# Process each expert
for expert_idx in range(self.num_experts):
count = counts_list[expert_idx]
if count == 0:
continue
start_idx = offsets_list[expert_idx]
end_idx = offsets_list[expert_idx + 1]
# Get tokens for this expert
expert_tokens = sorted_hidden[start_idx:end_idx]
expert_weights = sorted_routing_weights[start_idx:end_idx]
token_indices = sorted_token_indices[start_idx:end_idx]
# Get this expert's weight matrices
gate_w = self.expert_gate_proj[expert_idx]
up_w = self.expert_up_proj[expert_idx]
down_w = self.expert_down_proj[expert_idx]
# SwiGLU computation: down(silu(gate(x)) * up(x))
gate_out = F.linear(expert_tokens, gate_w)
up_out = F.linear(expert_tokens, up_w)
activated = F.silu(gate_out) * up_out
expert_out = F.linear(activated, down_w)
# Weight and accumulate
weighted_out = expert_out * expert_weights.unsqueeze(-1)
final_hidden_states.index_add_(0, token_indices, weighted_out)
return final_hidden_states
class _ScatterbrainRMSNormFallback(nn.Module):
"""Fallback RMSNorm when Liger is not available."""
def __init__(self, hidden_size, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(hidden_size))
self.variance_epsilon = eps
def forward(self, hidden_states):
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}"
# Use Liger RMSNorm if available, otherwise fallback
ScatterbrainRMSNorm = LigerRMSNorm if LIGER_AVAILABLE else _ScatterbrainRMSNormFallback
class ScatterbrainDecoderLayer(GradientCheckpointingLayer):
def __init__(self, config: ScatterbrainConfig, layer_idx: int):
super().__init__()
self.hidden_size = config.hidden_size
self.self_attn = ScatterbrainAttention(config, layer_idx)
# MoE layer
if config.num_experts > 0:
self.mlp = ScatterbrainSparseMoeBlock(config)
else:
self.mlp = ScatterbrainMLP(config, intermediate_size=config.intermediate_size)
self.input_layernorm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
@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] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values: Optional[Cache] = None,
cache_position: Optional[torch.LongTensor] = None,
cache_slot_idx: Optional[int] = None,
**kwargs: Unpack[FlashAttentionKwargs],
) -> torch.FloatTensor:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
# Self Attention
hidden_states, _ = self.self_attn(
hidden_states=hidden_states,
position_embeddings=position_embeddings,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
cache_position=cache_position,
cache_slot_idx=cache_slot_idx,
**kwargs,
)
hidden_states = residual + hidden_states
# MLP / MoE
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
# Unpack MoE output
if isinstance(hidden_states, tuple):
hidden_states, _ = hidden_states
hidden_states = residual + hidden_states
return hidden_states
class ScatterbrainRotaryEmbedding(nn.Module):
inv_freq: torch.Tensor
def __init__(self, config: ScatterbrainConfig, device=None):
super().__init__()
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"))
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
self.rope_init_fn = ROPE_INIT_FUNCTIONS[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
@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)
@auto_docstring
class ScatterbrainPreTrainedModel(PreTrainedModel):
config_class = ScatterbrainConfig
base_model_prefix = "model"
supports_gradient_checkpointing = True
_no_split_modules = ["ScatterbrainDecoderLayer"]
_skip_keys_device_placement = ["past_key_values"]
_supports_flash_attn = True
_supports_sdpa = True
_supports_flex_attn = True
_can_compile_fullgraph = False # MoE models don't work with torch.compile
_supports_attention_backend = True
@auto_docstring
class ScatterbrainModel(ScatterbrainPreTrainedModel):
def __init__(self, config: ScatterbrainConfig):
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)
# Only one physical layer
self.layers = nn.ModuleList(
[ScatterbrainDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
)
self.norm = ScatterbrainRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = ScatterbrainRotaryEmbedding(config=config)
self.gradient_checkpointing = False
# Number of loop iterations and cache slots
self._num_loop_iterations = config.num_loop_iterations
self._num_cache_slots = config.num_loop_iterations
# Initialize weights and apply final processing
self.post_init()
@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],
) -> MoeModelOutputWithPast:
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)
# Handle cache creation/extension for looping
if use_cache:
if past_key_values is None:
# Create cache with enough slots for all loop iterations
cache_config = copy.copy(self.config)
cache_config.num_hidden_layers = self._num_cache_slots
past_key_values = DynamicCache(config=cache_config)
elif isinstance(past_key_values, DynamicCache) and len(past_key_values.layers) < self._num_cache_slots:
# Extend cache if created externally with fewer slots
while len(past_key_values.layers) < self._num_cache_slots:
past_key_values.layers.append(DynamicLayer())
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)
mask_function = create_causal_mask if self.config.sliding_window is None else create_sliding_window_causal_mask
causal_mask = mask_function(
config=self.config,
input_embeds=inputs_embeds,
attention_mask=attention_mask,
cache_position=cache_position,
past_key_values=past_key_values,
position_ids=position_ids,
)
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids)
# Loop through the single layer multiple times
# Each iteration gets its own cache slot
decoder_layer = self.layers[0]
for loop_idx in range(self._num_loop_iterations):
hidden_states = decoder_layer(
hidden_states,
position_embeddings=position_embeddings,
attention_mask=causal_mask,
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
cache_position=cache_position,
cache_slot_idx=loop_idx,
**kwargs,
)
hidden_states = self.norm(hidden_states)
return MoeModelOutputWithPast(
last_hidden_state=hidden_states,
past_key_values=past_key_values,
)
@auto_docstring
class ScatterbrainForCausalLM(ScatterbrainPreTrainedModel, GenerationMixin):
_tied_weights_keys = ["lm_head.weight"]
def __init__(self, config):
super().__init__(config)
self.model = ScatterbrainModel(config)
self.vocab_size = config.vocab_size
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
self.router_aux_loss_coef = config.router_aux_loss_coef
self.num_experts = config.num_experts
self.num_experts_per_tok = config.num_experts_per_tok
# Initialize weights and apply final processing
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_embeddings):
self.lm_head = new_embeddings
@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,
output_router_logits: Optional[bool] = None,
cache_position: Optional[torch.LongTensor] = None,
logits_to_keep: Union[int, torch.Tensor] = 0,
**kwargs: Unpack[TransformersKwargs],
) -> MoeCausalLMOutputWithPast:
output_router_logits = (
output_router_logits if output_router_logits is not None else self.config.output_router_logits
)
outputs: MoeModelOutputWithPast = 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,
output_router_logits=output_router_logits,
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, labels, self.vocab_size, **kwargs)
return MoeCausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=outputs.past_key_values,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
router_logits=outputs.router_logits,
)
__all__ = [
"ScatterbrainConfig",
"ScatterbrainForCausalLM",
"ScatterbrainModel",
"ScatterbrainPreTrainedModel",
]