Text Generation
Transformers
Safetensors
taonet
trust-remote-code
sentencepiece
custom-architecture
custom_code
Instructions to use TaoTern/TaoNet-mini-A2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use TaoTern/TaoNet-mini-A2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="TaoTern/TaoNet-mini-A2", trust_remote_code=True)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("TaoTern/TaoNet-mini-A2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use TaoTern/TaoNet-mini-A2 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "TaoTern/TaoNet-mini-A2" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "TaoTern/TaoNet-mini-A2", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }'Use Docker
docker model run hf.co/TaoTern/TaoNet-mini-A2
- SGLang
How to use TaoTern/TaoNet-mini-A2 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 "TaoTern/TaoNet-mini-A2" \ --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": "TaoTern/TaoNet-mini-A2", "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 "TaoTern/TaoNet-mini-A2" \ --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": "TaoTern/TaoNet-mini-A2", "prompt": "Once upon a time,", "max_tokens": 512, "temperature": 0.5 }' - Docker Model Runner
How to use TaoTern/TaoNet-mini-A2 with Docker Model Runner:
docker model run hf.co/TaoTern/TaoNet-mini-A2
| """ | |
| SimpleLLM - Pure Attention-based Language Model with DeepSeek MLA + RoPE. | |
| Architecture: | |
| - Token Embedding → Attention Blocks → Output Head | |
| - Attention Blocks: Multi-head Latent Attention with RoPE positional embeddings | |
| - Feed-forward: SwiGLU gates | |
| - No state-space models (SSM), pure transformer architecture | |
| - Full BF16 precision (no quantization) | |
| """ | |
| import math | |
| from typing import Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from taoTrain.core import BaseModel | |
| from taoTrain.config import ModelConfig | |
| from .registry import register_architecture | |
| from .mla_components import AttentionBlock | |
| from .embeddings import FactorizedEmbedding | |
| class SimpleLLM(BaseModel): | |
| """ | |
| Pure attention-based language model with DeepSeek MLA + RoPE. | |
| Stateless architecture - no internal state management needed. | |
| Args: | |
| config: ModelConfig with: | |
| - vocab_size: Vocabulary size | |
| - hidden_dim: Model dimension (d_model) | |
| - hidden_dim_ff: Feed-forward dimension (default: 4 * hidden_dim) | |
| - num_layers: Number of attention blocks (n_layers) | |
| - num_heads: Number of attention heads (n_attn_heads) | |
| - d_latent_kv: KV compression dimension (default: 3/4 * hidden_dim) | |
| - d_rope: RoPE dimension per head (default: hidden_dim // num_heads) | |
| - max_seq_length: Maximum sequence length | |
| - dropout: Dropout rate | |
| - gqa_groups: Grouped Query Attention groups (default: 1) | |
| - use_factorized_embedding: Use low-rank embedding (default: False) | |
| """ | |
| def __init__(self, config: ModelConfig): | |
| super().__init__(config) | |
| # Parse config - use defaults if not specified | |
| self.vocab_size = config.vocab_size | |
| self.d_model = config.hidden_dim | |
| self.n_layers = config.num_layers | |
| self.n_heads = config.num_heads | |
| self.dropout = config.dropout | |
| # Optional parameters with smart defaults | |
| self.d_latent_kv = config.d_latent_kv if config.d_latent_kv is not None else int(self.d_model * 0.75) | |
| self.d_rope = config.d_rope if config.d_rope is not None else (self.d_model // self.n_heads) | |
| self.d_ff = config.hidden_dim_ff if config.hidden_dim_ff is not None else (self.d_model * 4) | |
| self.gqa_groups = getattr(config, 'gqa_groups', 1) | |
| self.use_factorized_embedding = getattr(config, 'use_factorized_embedding', False) | |
| self.d_embed_rank = getattr(config, 'd_embed_rank', 96) | |
| # YaRN parameters for context length extension | |
| self.rope_scale = getattr(config, 'rope_scale', 40.0) | |
| self.yarn_enabled = getattr(config, 'yarn_enabled', False) | |
| self.yarn_original_max_seq_length = getattr(config, 'yarn_original_max_seq_length', None) | |
| self.yarn_alpha = getattr(config, 'yarn_alpha', 1.0) | |
| self.max_seq_length = config.max_seq_length | |
| # Validate dimensions | |
| assert self.d_model % self.n_heads == 0, \ | |
| f"hidden_dim ({self.d_model}) must be divisible by num_heads ({self.n_heads})" | |
| assert self.d_latent_kv % self.n_heads == 0, \ | |
| f"d_latent_kv ({self.d_latent_kv}) must be divisible by num_heads ({self.n_heads})" | |
| # Token embedding | |
| if self.use_factorized_embedding: | |
| self.token_embedding = FactorizedEmbedding( | |
| self.vocab_size, | |
| self.d_model, | |
| self.d_embed_rank | |
| ) | |
| else: | |
| self.token_embedding = nn.Embedding(self.vocab_size, self.d_model) | |
| # Embedding dropout | |
| self.embedding_dropout = nn.Dropout(self.dropout) | |
| # Attention blocks with MLA + SwiGLU FFN | |
| self.blocks = nn.ModuleList() | |
| for _ in range(self.n_layers): | |
| self.blocks.append( | |
| AttentionBlock( | |
| d_model=self.d_model, | |
| d_latent_kv=self.d_latent_kv, | |
| n_heads=self.n_heads, | |
| d_rope=self.d_rope, | |
| d_ff=int(self.d_ff), | |
| dropout=self.dropout, | |
| gqa_groups=self.gqa_groups, | |
| rope_scale=self.rope_scale, | |
| max_seq_length=self.max_seq_length, | |
| yarn_enabled=self.yarn_enabled, | |
| yarn_original_max_seq_length=self.yarn_original_max_seq_length, | |
| yarn_alpha=self.yarn_alpha, | |
| ) | |
| ) | |
| # Final layer norm | |
| self.final_norm = nn.LayerNorm(self.d_model) | |
| # Output projection to vocabulary | |
| self.output_head = nn.Linear(self.d_model, self.vocab_size, bias=False) | |
| # Initialize weights | |
| self.apply(self._init_weights) | |
| # Cache for causal mask | |
| self.register_buffer("causal_mask_cache", None, persistent=False) | |
| self._print_architecture() | |
| def _init_weights(self, module): | |
| """Initialize weights for stable training.""" | |
| if isinstance(module, nn.Linear): | |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| if module.bias is not None: | |
| nn.init.zeros_(module.bias) | |
| elif isinstance(module, nn.Embedding): | |
| nn.init.normal_(module.weight, mean=0.0, std=0.02) | |
| def _print_architecture(self): | |
| """Print model architecture summary.""" | |
| total_params = sum(p.numel() for p in self.parameters()) | |
| trainable_params = sum(p.numel() for p in self.parameters() if p.requires_grad) | |
| print(f"\n{'='*70}") | |
| print("MODEL ARCHITECTURE - TAONET (DeepSeek MLA + RoPE)") | |
| print(f"{'='*70}") | |
| print(f"Embedding:") | |
| if self.use_factorized_embedding: | |
| embed_rank_params = self.vocab_size * self.d_embed_rank | |
| embed_proj_params = self.d_embed_rank * self.d_model | |
| print(f" Type: Factorized (rank={self.d_embed_rank})") | |
| print(f" Rank layer: {embed_rank_params/1e6:>8.2f}M") | |
| print(f" Projection: {embed_proj_params/1e6:>8.2f}M") | |
| else: | |
| embed_params = self.vocab_size * self.d_model | |
| print(f" Type: Standard") | |
| print(f" Params: {embed_params/1e6:>8.2f}M") | |
| output_params = self.d_model * self.vocab_size | |
| print(f"Output Head: {output_params/1e6:>8.2f}M") | |
| print(f"Attention Blocks: {len(self.blocks):>10} layers x AttentionBlock") | |
| print(f"{'-'*70}") | |
| print(f"Total Parameters: {total_params/1e6:>8.2f}M (trainable: {trainable_params/1e6:.2f}M)") | |
| print(f"{'-'*70}") | |
| print(f"Configuration:") | |
| print(f" Model dimension (d_model): {self.d_model}") | |
| print(f" KV latent dimension (d_latent_kv): {self.d_latent_kv}") | |
| print(f" Attention heads: {self.n_heads}") | |
| print(f" Head dimension: {self.d_model // self.n_heads}") | |
| print(f" RoPE dimension: {self.d_rope}") | |
| print(f" Feed-forward dimension: {int(self.d_ff)}") | |
| print(f" Number of layers: {self.n_layers}") | |
| print(f" Max sequence length: {self.config.max_seq_length}") | |
| print(f" Dropout: {self.dropout}") | |
| print(f" GQA groups: {self.gqa_groups}") | |
| print(f"{'='*70}\n") | |
| def _get_causal_mask(self, seq_len, device): | |
| """Get or create causal mask for sequence.""" | |
| if self.causal_mask_cache is None or self.causal_mask_cache.size(-1) < seq_len: | |
| # [seq_len, seq_len] lower triangular matrix (1 = attend, 0 = mask) | |
| mask = torch.tril(torch.ones(seq_len, seq_len, device=device, dtype=torch.bool)) | |
| self.register_buffer("causal_mask_cache", mask, persistent=False) | |
| return self.causal_mask_cache[:seq_len, :seq_len] | |
| def forward( | |
| self, | |
| input_ids: Optional[torch.Tensor] = None, | |
| attention_mask: Optional[torch.Tensor] = None, | |
| labels: Optional[torch.Tensor] = None, | |
| inputs_embeds: Optional[torch.Tensor] = None, | |
| pixel_values: Optional[torch.Tensor] = None, | |
| ) -> dict: | |
| """ | |
| Forward pass through the model. | |
| Args: | |
| input_ids: [batch_size, seq_len] tensor of token IDs | |
| attention_mask: [batch_size, seq_len] tensor where 1 = valid, 0 = padding | |
| labels: [batch_size, seq_len] target token IDs for loss computation | |
| Returns: | |
| Dictionary with: | |
| - 'logits': [batch_size, seq_len, vocab_size] output logits | |
| - 'loss': scalar loss (if labels provided, else None) | |
| """ | |
| if inputs_embeds is None: | |
| if input_ids is None: | |
| raise ValueError("Either input_ids or inputs_embeds must be provided") | |
| batch_size, seq_len = input_ids.shape | |
| device = input_ids.device | |
| x = self.token_embedding(input_ids) | |
| else: | |
| batch_size, seq_len, _ = inputs_embeds.shape | |
| device = inputs_embeds.device | |
| x = inputs_embeds | |
| # Get causal mask: [seq_len, seq_len] | |
| causal_mask = self._get_causal_mask(seq_len, device) | |
| # Combine causal mask with attention mask if provided | |
| if attention_mask is not None: | |
| # attention_mask: [batch, seq_len] where 1 = valid, 0 = padding | |
| # Expand to [batch, 1, 1, seq_len] | |
| padding_mask = attention_mask.unsqueeze(1).unsqueeze(1).bool() | |
| # Combine with causal: [1, 1, seq_len, seq_len] * [batch, 1, 1, seq_len] | |
| combined_mask = causal_mask.unsqueeze(0).unsqueeze(0) & padding_mask | |
| # For MLA: convert to {0, 1} format | |
| combined_mask = combined_mask.float() | |
| else: | |
| # Just causal mask | |
| combined_mask = causal_mask.unsqueeze(0).unsqueeze(0).float() | |
| x = self.embedding_dropout(x) | |
| # Pass through attention blocks | |
| for block in self.blocks: | |
| x = block(x, attention_mask=combined_mask) | |
| # Final layer norm | |
| x = self.final_norm(x) | |
| # Output projection to vocabulary | |
| logits = self.output_head(x) # [batch_size, seq_len, vocab_size] | |
| # Compute loss if labels are provided | |
| loss = None | |
| if labels is not None: | |
| # Flatten for loss computation | |
| logits_flat = logits.view(-1, logits.size(-1)) # (batch * seq_len, vocab_size) | |
| labels_flat = labels.view(-1) | |
| valid_label_mask = labels_flat != -100 | |
| if not torch.any(valid_label_mask): | |
| raise ValueError( | |
| "All labels are masked out (-100), so loss cannot be computed. " | |
| "This usually indicates a dataset parsing or masking bug." | |
| ) | |
| # Only compute loss on valid targets (ignore -100 tokens for padding) | |
| loss = F.cross_entropy( | |
| logits_flat, | |
| labels_flat, | |
| reduction='mean', | |
| ignore_index=-100 | |
| ) | |
| return { | |
| 'logits': logits, | |
| 'loss': loss, | |
| } | |