Text Generation
Transformers
Safetensors
English
slmoe
causal-lm
base-model
mixture-of-experts
sequence-routing
custom-code
trust-remote-code
custom_code
Instructions to use Banaxi-Tech/slmoe-test with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use Banaxi-Tech/slmoe-test with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="Banaxi-Tech/slmoe-test", trust_remote_code=True)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("Banaxi-Tech/slmoe-test", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use Banaxi-Tech/slmoe-test with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "Banaxi-Tech/slmoe-test" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Banaxi-Tech/slmoe-test", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/Banaxi-Tech/slmoe-test
- SGLang
How to use Banaxi-Tech/slmoe-test with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "Banaxi-Tech/slmoe-test" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Banaxi-Tech/slmoe-test", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "Banaxi-Tech/slmoe-test" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "Banaxi-Tech/slmoe-test", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use Banaxi-Tech/slmoe-test with Docker Model Runner:
docker model run hf.co/Banaxi-Tech/slmoe-test
| """Sequence-routed mixture-of-experts causal language model.""" | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import PreTrainedModel | |
| from transformers.cache_utils import Cache, DynamicCache | |
| from transformers.generation.utils import GenerationMixin | |
| from transformers.utils import ModelOutput | |
| try: | |
| from .configuration_slmoe import SLMoEConfig | |
| except ImportError: # Allows the standalone training script to import local code. | |
| from configuration_slmoe import SLMoEConfig | |
| class SLMoECausalLMOutputWithPast(ModelOutput): | |
| loss: Optional[torch.Tensor] = None | |
| logits: Optional[torch.Tensor] = None | |
| past_key_values: Optional[Cache] = None | |
| router_aux_loss: Optional[torch.Tensor] = None | |
| router_z_loss: Optional[torch.Tensor] = None | |
| expert_indices: Optional[torch.LongTensor] = None | |
| expert_weights: Optional[torch.Tensor] = None | |
| class SLMoERMSNorm(nn.Module): | |
| def __init__(self, dim: int, eps: float = 1e-6): | |
| super().__init__() | |
| self.eps = eps | |
| self.weight = nn.Parameter(torch.ones(dim)) | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| states = hidden_states.float() | |
| states = states * torch.rsqrt(states.square().mean(-1, keepdim=True) + self.eps) | |
| return (states * self.weight.float()).to(hidden_states.dtype) | |
| def _rope_cos_sin( | |
| head_dim: int, | |
| positions: torch.Tensor, | |
| theta: float, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| inv_freq = 1.0 / ( | |
| theta | |
| ** ( | |
| torch.arange(0, head_dim, 2, dtype=torch.float32, device=positions.device) | |
| / head_dim | |
| ) | |
| ) | |
| frequencies = torch.outer(positions.float(), inv_freq) | |
| return frequencies.cos(), frequencies.sin() | |
| def _apply_rope( | |
| query: torch.Tensor, | |
| key: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| query_dtype = query.dtype | |
| key_dtype = key.dtype | |
| cosine = cosine[None, None, :, :] | |
| sine = sine[None, None, :, :] | |
| query_pairs = query.float().reshape(*query.shape[:-1], -1, 2) | |
| key_pairs = key.float().reshape(*key.shape[:-1], -1, 2) | |
| query_even, query_odd = query_pairs.unbind(-1) | |
| key_even, key_odd = key_pairs.unbind(-1) | |
| query = torch.stack( | |
| (query_even * cosine - query_odd * sine, query_even * sine + query_odd * cosine), | |
| dim=-1, | |
| ).flatten(-2) | |
| key = torch.stack( | |
| (key_even * cosine - key_odd * sine, key_even * sine + key_odd * cosine), | |
| dim=-1, | |
| ).flatten(-2) | |
| return query.to(query_dtype), key.to(key_dtype) | |
| class SLMoECache(DynamicCache): | |
| """K/V cache carrying the one routing decision for the whole response.""" | |
| def __init__(self, config: SLMoEConfig): | |
| try: | |
| super().__init__(config=config) | |
| except TypeError: | |
| super().__init__() | |
| self.expert_indices: torch.LongTensor | None = None | |
| self.expert_weights: torch.Tensor | None = None | |
| def set_routing( | |
| self, | |
| expert_indices: torch.LongTensor, | |
| expert_weights: torch.Tensor, | |
| ) -> None: | |
| if self.expert_indices is not None: | |
| raise RuntimeError("The sequence routing plan may only be set once") | |
| self.expert_indices = expert_indices | |
| self.expert_weights = expert_weights | |
| def reorder_cache(self, beam_idx: torch.LongTensor): | |
| super().reorder_cache(beam_idx) | |
| if self.expert_indices is not None: | |
| beam_idx = beam_idx.to(self.expert_indices.device) | |
| self.expert_indices = self.expert_indices.index_select(0, beam_idx) | |
| self.expert_weights = self.expert_weights.index_select(0, beam_idx) | |
| def batch_repeat_interleave(self, repeats: int): | |
| super().batch_repeat_interleave(repeats) | |
| if self.expert_indices is not None: | |
| self.expert_indices = self.expert_indices.repeat_interleave(repeats, dim=0) | |
| self.expert_weights = self.expert_weights.repeat_interleave(repeats, dim=0) | |
| def batch_select_indices(self, indices: torch.Tensor): | |
| super().batch_select_indices(indices) | |
| if self.expert_indices is not None: | |
| indices = indices.to(self.expert_indices.device) | |
| self.expert_indices = self.expert_indices.index_select(0, indices) | |
| self.expert_weights = self.expert_weights.index_select(0, indices) | |
| class SLMoEAttention(nn.Module): | |
| def __init__(self, config: SLMoEConfig, layer_idx: int): | |
| super().__init__() | |
| self.layer_idx = layer_idx | |
| self.num_heads = config.num_attention_heads | |
| self.num_kv_heads = config.num_key_value_heads | |
| self.head_dim = config.head_dim | |
| self.num_kv_groups = self.num_heads // self.num_kv_heads | |
| self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=False) | |
| self.k_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False) | |
| self.v_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False) | |
| self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=False) | |
| self.o_proj.SLMOE_SCALE_INIT = True | |
| self.q_norm = SLMoERMSNorm(self.head_dim, config.rms_norm_eps) | |
| self.k_norm = SLMoERMSNorm(self.head_dim, config.rms_norm_eps) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| attention_mask: torch.Tensor | None = None, | |
| past_key_values: Cache | None = None, | |
| ) -> torch.Tensor: | |
| batch_size, query_length, _ = hidden_states.shape | |
| query = self.q_proj(hidden_states).view( | |
| batch_size, query_length, self.num_heads, self.head_dim | |
| ).transpose(1, 2) | |
| key = self.k_proj(hidden_states).view( | |
| batch_size, query_length, self.num_kv_heads, self.head_dim | |
| ).transpose(1, 2) | |
| value = self.v_proj(hidden_states).view( | |
| batch_size, query_length, self.num_kv_heads, self.head_dim | |
| ).transpose(1, 2) | |
| query = self.q_norm(query) | |
| key = self.k_norm(key) | |
| query, key = _apply_rope(query, key, cosine, sine) | |
| past_length = 0 | |
| if past_key_values is not None: | |
| past_length = past_key_values.get_seq_length(self.layer_idx) | |
| key, value = past_key_values.update(key, value, self.layer_idx) | |
| key_length = key.size(-2) | |
| key = key.repeat_interleave(self.num_kv_groups, dim=1) | |
| value = value.repeat_interleave(self.num_kv_groups, dim=1) | |
| is_causal = query_length > 1 and past_length == 0 and attention_mask is None | |
| sdpa_mask = None | |
| if not is_causal and query_length > 1: | |
| query_positions = past_length + torch.arange(query_length, device=query.device) | |
| key_positions = torch.arange(key_length, device=query.device) | |
| sdpa_mask = (key_positions[None, :] <= query_positions[:, None])[None, None] | |
| if attention_mask is not None: | |
| key_padding = attention_mask.to(torch.bool) | |
| if key_padding.size(-1) < key_length: | |
| key_padding = F.pad(key_padding, (key_length - key_padding.size(-1), 0), value=True) | |
| else: | |
| key_padding = key_padding[:, -key_length:] | |
| key_padding = key_padding[:, None, None, :] | |
| sdpa_mask = key_padding if sdpa_mask is None else sdpa_mask & key_padding | |
| is_causal = False | |
| output = F.scaled_dot_product_attention( | |
| query, | |
| key, | |
| value, | |
| attn_mask=sdpa_mask, | |
| is_causal=is_causal, | |
| ) | |
| output = output.transpose(1, 2).contiguous().view( | |
| batch_size, query_length, self.num_heads * self.head_dim | |
| ) | |
| return self.o_proj(output) | |
| class SLMoESequenceRouter(nn.Module): | |
| """Choose one fixed expert set from a causal prefix of each sequence.""" | |
| def __init__(self, config: SLMoEConfig): | |
| super().__init__() | |
| self.num_experts = config.num_experts | |
| self.top_k = config.num_experts_per_sequence | |
| self.prefix_length = config.router_prefix_length | |
| self.jitter_noise = config.router_jitter_noise | |
| self.norm = SLMoERMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.proj = nn.Linear(config.hidden_size, config.num_experts, bias=False) | |
| def prefix_mask( | |
| self, | |
| token_embeddings: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| batch_size, sequence_length, _ = token_embeddings.shape | |
| if attention_mask is None: | |
| positions = torch.arange(sequence_length, device=token_embeddings.device) | |
| return (positions < self.prefix_length).expand(batch_size, -1) | |
| valid = attention_mask[:, -sequence_length:].to(torch.bool) | |
| valid_order = valid.long().cumsum(dim=-1) | |
| return valid & (valid_order <= self.prefix_length) | |
| def forward( | |
| self, | |
| token_embeddings: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| ) -> tuple[torch.LongTensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| prefix_mask = self.prefix_mask(token_embeddings, attention_mask) | |
| normalized = self.norm(token_embeddings) | |
| mask = prefix_mask.unsqueeze(-1).to(normalized.dtype) | |
| pooled = (normalized * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0) | |
| if self.training and self.jitter_noise > 0: | |
| pooled = pooled * torch.empty_like(pooled).uniform_( | |
| 1.0 - self.jitter_noise, | |
| 1.0 + self.jitter_noise, | |
| ) | |
| router_logits = self.proj(pooled).float() | |
| router_probs = F.softmax(router_logits, dim=-1, dtype=torch.float32) | |
| top_probs, expert_indices = torch.topk( | |
| router_probs, | |
| k=self.top_k, | |
| dim=-1, | |
| sorted=True, | |
| ) | |
| expert_weights = top_probs / top_probs.sum(dim=-1, keepdim=True).clamp_min(1e-9) | |
| selected_fraction = F.one_hot( | |
| expert_indices, | |
| num_classes=self.num_experts, | |
| ).float().mean(dim=(0, 1)) | |
| probability_fraction = router_probs.mean(dim=0) | |
| auxiliary_loss = self.num_experts * torch.sum( | |
| selected_fraction * probability_fraction | |
| ) | |
| router_z_loss = torch.logsumexp(router_logits, dim=-1).square().mean() | |
| return expert_indices, expert_weights, auxiliary_loss, router_z_loss | |
| class SLMoEExpertBank(nn.Module): | |
| """Batched expert weights; only the sequence-selected slices are evaluated.""" | |
| def __init__(self, config: SLMoEConfig): | |
| super().__init__() | |
| experts = config.num_experts | |
| hidden = config.hidden_size | |
| intermediate = config.expert_intermediate_size | |
| self.output_scale = config.expert_output_scale | |
| self.gate_weight = nn.Parameter(torch.empty(experts, intermediate, hidden)) | |
| self.up_weight = nn.Parameter(torch.empty(experts, intermediate, hidden)) | |
| self.down_weight = nn.Parameter(torch.empty(experts, hidden, intermediate)) | |
| nn.init.normal_(self.gate_weight, mean=0.0, std=config.initializer_range) | |
| nn.init.normal_(self.up_weight, mean=0.0, std=config.initializer_range) | |
| down_std = config.initializer_range * (2 * config.num_hidden_layers) ** -0.5 | |
| nn.init.normal_(self.down_weight, mean=0.0, std=down_std) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| expert_indices: torch.LongTensor, | |
| expert_weights: torch.Tensor, | |
| ) -> torch.Tensor: | |
| gate_weight = self.gate_weight[expert_indices] | |
| up_weight = self.up_weight[expert_indices] | |
| down_weight = self.down_weight[expert_indices] | |
| gate = torch.einsum("bsh,bkih->bski", hidden_states, gate_weight) | |
| up = torch.einsum("bsh,bkih->bski", hidden_states, up_weight) | |
| activated = F.silu(gate) * up | |
| activated = activated * expert_weights[:, None, :, None].to(activated.dtype) | |
| output = torch.einsum("bski,bkhi->bsh", activated, down_weight) | |
| return output * self.output_scale | |
| class SLMoEBlock(nn.Module): | |
| def __init__(self, config: SLMoEConfig, layer_idx: int): | |
| super().__init__() | |
| self.input_norm = SLMoERMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.attention = SLMoEAttention(config, layer_idx) | |
| self.post_attention_norm = SLMoERMSNorm(config.hidden_size, config.rms_norm_eps) | |
| self.experts = SLMoEExpertBank(config) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| cosine: torch.Tensor, | |
| sine: torch.Tensor, | |
| expert_indices: torch.LongTensor, | |
| expert_weights: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| past_key_values: Cache | None, | |
| ) -> torch.Tensor: | |
| hidden_states = hidden_states + self.attention( | |
| self.input_norm(hidden_states), | |
| cosine, | |
| sine, | |
| attention_mask=attention_mask, | |
| past_key_values=past_key_values, | |
| ) | |
| return hidden_states + self.experts( | |
| self.post_attention_norm(hidden_states), | |
| expert_indices, | |
| expert_weights, | |
| ) | |
| class SLMoEPreTrainedModel(PreTrainedModel): | |
| config_class = SLMoEConfig | |
| base_model_prefix = "transformer" | |
| supports_gradient_checkpointing = False | |
| _no_split_modules = ["SLMoEBlock"] | |
| _supports_sdpa = True | |
| _supports_cache_class = True | |
| def _init_weights(self, module: nn.Module): | |
| std = self.config.initializer_range | |
| if hasattr(module, "SLMOE_SCALE_INIT"): | |
| std *= (2 * self.config.num_hidden_layers) ** -0.5 | |
| if isinstance(module, nn.Linear): | |
| nn.init.normal_(module.weight, mean=0.0, std=std) | |
| elif isinstance(module, nn.Embedding): | |
| nn.init.normal_(module.weight, mean=0.0, std=std) | |
| class SLMoEForCausalLM(SLMoEPreTrainedModel, GenerationMixin): | |
| _tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"} | |
| def _supports_default_dynamic_cache(cls) -> bool: | |
| return False | |
| def __init__(self, config: SLMoEConfig): | |
| super().__init__(config) | |
| self.router = SLMoESequenceRouter(config) | |
| self.transformer = nn.ModuleDict( | |
| { | |
| "wte": nn.Embedding(config.vocab_size, config.hidden_size), | |
| "h": nn.ModuleList( | |
| [SLMoEBlock(config, index) for index in range(config.num_hidden_layers)] | |
| ), | |
| "ln_f": SLMoERMSNorm(config.hidden_size, config.rms_norm_eps), | |
| } | |
| ) | |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) | |
| self.embedding_scale = math.sqrt(config.hidden_size) | |
| self.post_init() | |
| if config.tie_word_embeddings: | |
| self.tie_weights() | |
| def get_input_embeddings(self): | |
| return self.transformer["wte"] | |
| def set_input_embeddings(self, value): | |
| self.transformer["wte"] = value | |
| def get_output_embeddings(self): | |
| return self.lm_head | |
| def set_output_embeddings(self, value): | |
| self.lm_head = value | |
| def _route( | |
| self, | |
| token_embeddings: torch.Tensor, | |
| attention_mask: torch.Tensor | None, | |
| past_key_values: SLMoECache | None, | |
| ) -> tuple[torch.LongTensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| if past_key_values is not None and past_key_values.expert_indices is not None: | |
| zero = token_embeddings.new_zeros((), dtype=torch.float32) | |
| return ( | |
| past_key_values.expert_indices, | |
| past_key_values.expert_weights, | |
| zero, | |
| zero, | |
| ) | |
| expert_indices, expert_weights, auxiliary_loss, router_z_loss = self.router( | |
| token_embeddings, | |
| attention_mask, | |
| ) | |
| if past_key_values is not None: | |
| past_key_values.set_routing(expert_indices, expert_weights) | |
| return expert_indices, expert_weights, auxiliary_loss, router_z_loss | |
| def forward( | |
| self, | |
| input_ids: torch.LongTensor, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| labels: Optional[torch.LongTensor] = None, | |
| past_key_values: Optional[Cache] = None, | |
| use_cache: Optional[bool] = None, | |
| **kwargs, | |
| ) -> SLMoECausalLMOutputWithPast: | |
| if use_cache is None: | |
| use_cache = self.config.use_cache and labels is None | |
| if use_cache and past_key_values is None: | |
| past_key_values = SLMoECache(self.config) | |
| if use_cache and not isinstance(past_key_values, SLMoECache): | |
| raise TypeError("SLMoE requires SLMoECache to preserve sequence routing") | |
| if not use_cache: | |
| past_key_values = None | |
| past_length = past_key_values.get_seq_length() if past_key_values is not None else 0 | |
| sequence_length = input_ids.size(1) | |
| total_length = past_length + sequence_length | |
| if total_length > self.config.max_position_embeddings: | |
| raise ValueError( | |
| f"Sequence length {total_length} exceeds {self.config.max_position_embeddings}" | |
| ) | |
| token_embeddings = self.transformer["wte"](input_ids) | |
| expert_indices, expert_weights, router_aux_loss, router_z_loss = self._route( | |
| token_embeddings, | |
| attention_mask, | |
| past_key_values, | |
| ) | |
| hidden_states = token_embeddings * self.embedding_scale | |
| positions = torch.arange( | |
| past_length, | |
| total_length, | |
| dtype=torch.float32, | |
| device=input_ids.device, | |
| ) | |
| cosine, sine = _rope_cos_sin( | |
| self.config.head_dim, | |
| positions, | |
| self.config.rope_theta, | |
| ) | |
| for block in self.transformer["h"]: | |
| hidden_states = block( | |
| hidden_states, | |
| cosine, | |
| sine, | |
| expert_indices, | |
| expert_weights, | |
| attention_mask, | |
| past_key_values, | |
| ) | |
| hidden_states = self.transformer["ln_f"](hidden_states) | |
| logits = self.lm_head(hidden_states) | |
| loss = None | |
| if labels is not None: | |
| shift_logits = logits[..., :-1, :].float().contiguous() | |
| shift_labels = labels[..., 1:].clone().contiguous() | |
| prefix_mask = self.router.prefix_mask(token_embeddings, attention_mask) | |
| sequence_positions = torch.arange(sequence_length, device=input_ids.device) | |
| last_prefix_position = torch.where( | |
| prefix_mask, | |
| sequence_positions[None, :], | |
| -1, | |
| ).amax(dim=-1) | |
| prediction_positions = sequence_positions[:-1][None, :] | |
| shift_labels[prediction_positions < last_prefix_position[:, None]] = -100 | |
| ce_loss = F.cross_entropy( | |
| shift_logits.reshape(-1, shift_logits.size(-1)), | |
| shift_labels.reshape(-1), | |
| ignore_index=-100, | |
| ) | |
| loss = ( | |
| ce_loss | |
| + self.config.router_aux_loss_coeff * router_aux_loss | |
| + self.config.router_z_loss_coeff * router_z_loss | |
| ) | |
| return SLMoECausalLMOutputWithPast( | |
| loss=loss, | |
| logits=logits, | |
| past_key_values=past_key_values, | |
| router_aux_loss=router_aux_loss, | |
| router_z_loss=router_z_loss, | |
| expert_indices=expert_indices, | |
| expert_weights=expert_weights, | |
| ) | |
| SLMoEForCausalLM.register_for_auto_class("AutoModelForCausalLM") | |
| __all__ = [ | |
| "SLMoECache", | |
| "SLMoECausalLMOutputWithPast", | |
| "SLMoEForCausalLM", | |
| "SLMoEPreTrainedModel", | |
| ] | |