# Transformers Library Integration Guide ## Overview This guide covers integrating custom CUDA kernels into HuggingFace Transformers models, focusing on LLaMA, Mistral, and Qwen architectures. While similar to diffusers integration in concept, transformers models have distinct patterns, class hierarchies, and conventions that require different handling. ## Model Architecture Analysis ### Inspecting a Transformers Model ```python from transformers import AutoModelForCausalLM, AutoConfig import torch # Load model model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-7b-hf", torch_dtype=torch.bfloat16, device_map="cuda", ) # List all module types from collections import Counter type_counts = Counter( type(m).__name__ for _, m in model.named_modules() ) print("Module types:") for name, count in type_counts.most_common(): print(f" {name}: {count}") # Example output for LLaMA-7B: # LlamaDecoderLayer: 32 # LlamaRMSNorm: 65 (2 per layer + 1 final) # LlamaAttention: 32 # LlamaMLP: 32 # LlamaRotaryEmbedding: 32 # Linear: 224 ``` ### Identifying Kernel Targets ```python # Profile to find hotspots from torch.profiler import profile, ProfilerActivity with torch.no_grad(): input_ids = torch.randint(0, 32000, (1, 512), device="cuda") with profile( activities=[ProfilerActivity.CUDA], record_shapes=True, ) as prof: model(input_ids) print(prof.key_averages().table( sort_by="cuda_time_total", row_limit=15, )) ``` ## LLaMA / Mistral / Qwen Architectures ### LLaMA Architecture ``` LlamaForCausalLM ├── model (LlamaModel) │ ├── embed_tokens # Token embedding │ ├── layers # LlamaDecoderLayer x N │ │ ├── input_layernorm # LlamaRMSNorm (pre-attention) │ │ ├── self_attn # LlamaAttention │ │ │ ├── q_proj # Linear │ │ │ ├── k_proj # Linear (GQA: fewer heads) │ │ │ ├── v_proj # Linear (GQA: fewer heads) │ │ │ ├── o_proj # Linear │ │ │ └── rotary_emb # LlamaRotaryEmbedding │ │ ├── post_attention_layernorm # LlamaRMSNorm (pre-MLP) │ │ └── mlp # LlamaMLP │ │ ├── gate_proj # Linear (SiLU gate) │ │ ├── up_proj # Linear │ │ └── down_proj # Linear │ └── norm # LlamaRMSNorm (final) └── lm_head # Linear (vocabulary projection) ``` ### Mistral Architecture Mistral is very similar to LLaMA but with key differences: ``` MistralForCausalLM ├── model (MistralModel) │ ├── embed_tokens │ ├── layers # MistralDecoderLayer x N │ │ ├── input_layernorm # MistralRMSNorm │ │ ├── self_attn # MistralAttention │ │ │ ├── q_proj, k_proj, v_proj, o_proj │ │ │ └── rotary_emb │ │ ├── post_attention_layernorm # MistralRMSNorm │ │ └── mlp # MistralMLP │ │ ├── gate_proj # SiLU gate │ │ ├── up_proj │ │ └── down_proj │ └── norm # MistralRMSNorm (final) └── lm_head ``` Key Mistral differences: - Uses **Sliding Window Attention** (window_size=4096 by default) - Uses **Grouped Query Attention** (GQA) like LLaMA 2 ### Qwen Architecture ``` Qwen2ForCausalLM ├── model (Qwen2Model) │ ├── embed_tokens │ ├── layers # Qwen2DecoderLayer x N │ │ ├── input_layernorm # Qwen2RMSNorm │ │ ├── self_attn # Qwen2Attention │ │ │ ├── q_proj, k_proj, v_proj, o_proj │ │ │ └── rotary_emb │ │ ├── post_attention_layernorm # Qwen2RMSNorm │ │ └── mlp # Qwen2MLP │ │ ├── gate_proj # SiLU gate │ │ ├── up_proj │ │ └── down_proj │ └── norm # Qwen2RMSNorm (final) └── lm_head ``` ### Architecture Comparison Table | Feature | LLaMA 2 | LLaMA 3 | Mistral 7B | Qwen 2 | |---|---|---|---|---| | RMSNorm class | LlamaRMSNorm | LlamaRMSNorm | MistralRMSNorm | Qwen2RMSNorm | | Epsilon attribute | `variance_epsilon` | `variance_epsilon` | `variance_epsilon` | `variance_epsilon` | | Activation | SiLU (gate) | SiLU (gate) | SiLU (gate) | SiLU (gate) | | Attention | GQA | GQA | GQA + Sliding Window | GQA | | RoPE | Standard | Extended (RoPE scaling) | Standard | Standard | | Hidden sizes | 4096, 5120 | 4096, 8192 | 4096 | 3584, 4096, 8192 | | Weight always present | **Yes** | **Yes** | **Yes** | **Yes** | | `set_processor` method | No | No | No | No | ## RMSNorm Patcher for Transformers ### Critical Difference: variance_epsilon vs eps In transformers models, the epsilon attribute is named `variance_epsilon`, **not** `eps` as in diffusers: ```python # Diffusers RMSNorm: from diffusers.models.normalization import RMSNorm as DiffusersRMSNorm # norm = DiffusersRMSNorm(hidden_size, eps=1e-6) # norm.eps <-- attribute name # Transformers LlamaRMSNorm: from transformers.models.llama.modeling_llama import LlamaRMSNorm # norm = LlamaRMSNorm(hidden_size, eps=1e-6) # norm.variance_epsilon <-- different attribute name! # Transformers MistralRMSNorm: from transformers.models.mistral.modeling_mistral import MistralRMSNorm # norm.variance_epsilon <-- same as LLaMA # Transformers Qwen2RMSNorm: from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm # norm.variance_epsilon <-- same as LLaMA ``` ### Universal RMSNorm Patcher ```python import torch import torch.nn as nn from typing import Optional, Dict, Type class TransformersRMSNormPatcher: """ Patches RMSNorm modules in transformers models. Handles the variance_epsilon vs eps naming difference automatically. """ # Known RMSNorm classes and their epsilon attribute names KNOWN_CLASSES: Dict[str, str] = {} def __init__(self, cuda_rmsnorm_fn): self.cuda_rmsnorm = cuda_rmsnorm_fn self._original_forwards = {} @classmethod def _discover_rmsnorm_classes(cls): """Discover available RMSNorm classes from transformers.""" rmsnorm_classes = {} # LLaMA try: from transformers.models.llama.modeling_llama import LlamaRMSNorm rmsnorm_classes[LlamaRMSNorm] = "variance_epsilon" except ImportError: pass # Mistral try: from transformers.models.mistral.modeling_mistral import MistralRMSNorm rmsnorm_classes[MistralRMSNorm] = "variance_epsilon" except ImportError: pass # Qwen2 try: from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm rmsnorm_classes[Qwen2RMSNorm] = "variance_epsilon" except ImportError: pass # Gemma try: from transformers.models.gemma.modeling_gemma import GemmaRMSNorm rmsnorm_classes[GemmaRMSNorm] = "eps" except ImportError: pass return rmsnorm_classes def _get_epsilon(self, module: nn.Module) -> float: """Get the epsilon value from an RMSNorm module, handling different attribute names.""" # Try common attribute names for attr_name in ["variance_epsilon", "eps", "epsilon", "rms_norm_eps"]: if hasattr(module, attr_name): return getattr(module, attr_name) # Default fallback return 1e-6 def patch(self, model: nn.Module) -> int: """ Patch all RMSNorm modules in a transformers model. Returns the number of modules patched. """ rmsnorm_classes = self._discover_rmsnorm_classes() count = 0 for name, module in model.named_modules(): # Check if module is an instance of any known RMSNorm class is_rmsnorm = any( isinstance(module, cls) for cls in rmsnorm_classes ) if is_rmsnorm: self._original_forwards[name] = module.forward eps = self._get_epsilon(module) module.forward = self._make_forward(module, eps) count += 1 return count def _make_forward(self, module: nn.Module, eps: float): """Create a patched forward function.""" cuda_fn = self.cuda_rmsnorm def patched_forward(hidden_states: torch.Tensor) -> torch.Tensor: # In transformers, weight is ALWAYS present (unlike diffusers) return cuda_fn(hidden_states, module.weight, eps) return patched_forward def unpatch(self, model: nn.Module) -> int: """Restore original forward methods.""" count = 0 for name, module in model.named_modules(): if name in self._original_forwards: module.forward = self._original_forwards[name] count += 1 self._original_forwards.clear() return count ``` ## Key Differences from Diffusers Understanding these differences is critical to avoid subtle bugs. ### 1. Weight is Always Present In transformers models, RMSNorm **always** has a weight parameter. You do not need to handle `weight=None`: ```python # Diffusers: weight can be None def diffusers_rmsnorm_forward(module, hidden_states): if module.weight is None: return rmsnorm_no_weight(hidden_states, module.eps) return rmsnorm(hidden_states, module.weight, module.eps) # Transformers: weight is ALWAYS present def transformers_rmsnorm_forward(module, hidden_states): # No None check needed -- weight always exists return rmsnorm(hidden_states, module.weight, module.variance_epsilon) ``` ### 2. No set_processor for Attention Transformers attention modules do **not** have a `set_processor` method like diffusers. You must patch the forward method directly: ```python # Diffusers: Use set_processor # module.set_processor(CustomProcessor()) # Transformers: Patch forward directly from transformers.models.llama.modeling_llama import LlamaAttention for name, module in model.named_modules(): if isinstance(module, LlamaAttention): original_forward = module.forward def make_forward(mod, orig_fn): def patched_forward(*args, **kwargs): # Custom logic before/after attention # Or completely replace attention computation return orig_fn(*args, **kwargs) return patched_forward module.forward = make_forward(module, original_forward) ``` ### 3. device_map Handling Transformers models support `device_map` for automatic model parallelism, which can place different layers on different devices. Your kernel injection must respect this: ```python # The model may be split across devices model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-70b-hf", device_map="auto", # Splits model across available GPUs torch_dtype=torch.bfloat16, ) # When patching, check which device each module is on def inject_kernels_with_device_map(model): for name, module in model.named_modules(): if isinstance(module, LlamaRMSNorm): # Check the device of this module's parameters device = next(module.parameters()).device if device.type == "cuda": # Patch with CUDA kernel module.forward = make_cuda_forward(module) else: # Module is on CPU -- skip patching print(f"Skipping {name}: on {device}") ``` ### 4. Model-Specific Import Paths Each model family has its own module path: ```python # LLaMA from transformers.models.llama.modeling_llama import ( LlamaRMSNorm, LlamaAttention, LlamaMLP, LlamaDecoderLayer, ) # Mistral from transformers.models.mistral.modeling_mistral import ( MistralRMSNorm, MistralAttention, MistralMLP, ) # Qwen2 from transformers.models.qwen2.modeling_qwen2 import ( Qwen2RMSNorm, Qwen2Attention, Qwen2MLP, ) # Generic approach: detect model type at runtime def get_rmsnorm_class(model): """Get the RMSNorm class used by this model.""" model_type = model.config.model_type class_map = { "llama": "transformers.models.llama.modeling_llama.LlamaRMSNorm", "mistral": "transformers.models.mistral.modeling_mistral.MistralRMSNorm", "qwen2": "transformers.models.qwen2.modeling_qwen2.Qwen2RMSNorm", } if model_type in class_map: module_path, class_name = class_map[model_type].rsplit(".", 1) import importlib mod = importlib.import_module(module_path) return getattr(mod, class_name) return None ``` ## Model-Specific Integration ### LLaMA Integration ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer from transformers.models.llama.modeling_llama import LlamaRMSNorm def inject_llama_kernels(model, cuda_rmsnorm_fn): """Inject custom CUDA kernels into a LLaMA model.""" count = 0 for name, module in model.named_modules(): if isinstance(module, LlamaRMSNorm): eps = module.variance_epsilon def make_forward(mod, e): def forward(hidden_states): return cuda_rmsnorm_fn(hidden_states, mod.weight, e) return forward module.forward = make_forward(module, eps) count += 1 print(f"Patched {count} LlamaRMSNorm modules") return model # Usage model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-7b-hf", torch_dtype=torch.bfloat16, device_map="cuda", ) inject_llama_kernels(model, cuda_rmsnorm) # Verify tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-hf") inputs = tokenizer("The capital of France is", return_tensors="pt").to("cuda") outputs = model.generate(**inputs, max_new_tokens=20) print(tokenizer.decode(outputs[0], skip_special_tokens=True)) ``` ### Mistral Integration ```python from transformers.models.mistral.modeling_mistral import MistralRMSNorm def inject_mistral_kernels(model, cuda_rmsnorm_fn): """Inject custom CUDA kernels into a Mistral model.""" count = 0 for name, module in model.named_modules(): if isinstance(module, MistralRMSNorm): eps = module.variance_epsilon def make_forward(mod, e): def forward(hidden_states): return cuda_rmsnorm_fn(hidden_states, mod.weight, e) return forward module.forward = make_forward(module, eps) count += 1 print(f"Patched {count} MistralRMSNorm modules") return model ``` ### Qwen Integration ```python from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm def inject_qwen_kernels(model, cuda_rmsnorm_fn): """Inject custom CUDA kernels into a Qwen model.""" count = 0 for name, module in model.named_modules(): if isinstance(module, Qwen2RMSNorm): eps = module.variance_epsilon def make_forward(mod, e): def forward(hidden_states): return cuda_rmsnorm_fn(hidden_states, mod.weight, e) return forward module.forward = make_forward(module, eps) count += 1 print(f"Patched {count} Qwen2RMSNorm modules") return model ``` ### Universal Integration Function ```python def inject_transformers_kernels(model, cuda_rmsnorm_fn): """ Universal kernel injection for any supported transformers model. Automatically detects model type and patches accordingly. """ model_type = model.config.model_type print(f"Detected model type: {model_type}") # Map model types to their RMSNorm classes rmsnorm_imports = { "llama": ("transformers.models.llama.modeling_llama", "LlamaRMSNorm"), "mistral": ("transformers.models.mistral.modeling_mistral", "MistralRMSNorm"), "qwen2": ("transformers.models.qwen2.modeling_qwen2", "Qwen2RMSNorm"), "gemma": ("transformers.models.gemma.modeling_gemma", "GemmaRMSNorm"), "phi3": ("transformers.models.phi3.modeling_phi3", "Phi3RMSNorm"), } if model_type not in rmsnorm_imports: print(f"Warning: Unknown model type '{model_type}'. Attempting generic patching.") return _generic_rmsnorm_patch(model, cuda_rmsnorm_fn) module_path, class_name = rmsnorm_imports[model_type] import importlib mod = importlib.import_module(module_path) rmsnorm_cls = getattr(mod, class_name) count = 0 for name, module in model.named_modules(): if isinstance(module, rmsnorm_cls): eps = _get_epsilon(module) def make_forward(m, e): def forward(hidden_states): return cuda_rmsnorm_fn(hidden_states, m.weight, e) return forward module.forward = make_forward(module, eps) count += 1 print(f"Patched {count} {class_name} modules") return model def _get_epsilon(module): """Extract epsilon from a module, handling various attribute names.""" for attr in ["variance_epsilon", "eps", "epsilon", "rms_norm_eps"]: if hasattr(module, attr): return getattr(module, attr) return 1e-6 def _generic_rmsnorm_patch(model, cuda_rmsnorm_fn): """Fallback: patch any module whose class name contains 'RMSNorm'.""" count = 0 for name, module in model.named_modules(): class_name = type(module).__name__ if "RMSNorm" in class_name and hasattr(module, "weight"): eps = _get_epsilon(module) def make_forward(m, e): def forward(hidden_states): return cuda_rmsnorm_fn(hidden_states, m.weight, e) return forward module.forward = make_forward(module, eps) count += 1 print(f"Generic patch: patched {count} RMSNorm-like modules") return model ``` ## Flash Attention 2 Integration Transformers supports Flash Attention 2 natively. Custom kernels should be compatible. ### Enabling Flash Attention 2 ```python from transformers import AutoModelForCausalLM # Method 1: Use attn_implementation parameter model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-7b-hf", torch_dtype=torch.bfloat16, attn_implementation="flash_attention_2", device_map="cuda", ) # Method 2: Use BetterTransformer (deprecated, prefer method 1) # model = model.to_bettertransformer() ``` ### Custom Attention with Flash Attention Fallback ```python import torch import torch.nn.functional as F def patched_attention_forward( self, hidden_states, attention_mask=None, position_ids=None, past_key_value=None, output_attentions=False, use_cache=False, **kwargs, ): """ Custom attention forward that uses custom RoPE kernel but falls back to Flash Attention for the attention computation. """ bsz, q_len, _ = hidden_states.size() # QKV projection (standard) query_states = self.q_proj(hidden_states) key_states = self.k_proj(hidden_states) value_states = self.v_proj(hidden_states) # Reshape for multi-head query_states = query_states.view( bsz, q_len, self.num_heads, self.head_dim ).transpose(1, 2) key_states = key_states.view( bsz, q_len, self.num_key_value_heads, self.head_dim ).transpose(1, 2) value_states = value_states.view( bsz, q_len, self.num_key_value_heads, self.head_dim ).transpose(1, 2) # Apply RoPE using custom CUDA kernel cos, sin = self.rotary_emb(value_states, position_ids) query_states = cuda_apply_rope(query_states, cos, sin) key_states = cuda_apply_rope(key_states, cos, sin) # KV cache handling if past_key_value is not None: key_states, value_states = past_key_value.update( key_states, value_states, self.layer_idx ) # GQA: repeat k/v heads key_states = repeat_kv(key_states, self.num_key_value_groups) value_states = repeat_kv(value_states, self.num_key_value_groups) # Use PyTorch SDPA (which calls Flash Attention on supported hardware) attn_output = F.scaled_dot_product_attention( query_states, key_states, value_states, attn_mask=attention_mask, dropout_p=0.0, is_causal=attention_mask is None and q_len > 1, ) # Output projection attn_output = attn_output.transpose(1, 2).contiguous() attn_output = attn_output.reshape(bsz, q_len, -1) attn_output = self.o_proj(attn_output) return attn_output, None, past_key_value ``` ## Verification ### Correctness Verification ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer def verify_transformers_correctness( model_name: str, inject_fn, cuda_rmsnorm_fn, max_new_tokens: int = 50, ): """ Verify that custom kernels produce identical outputs for a transformers model. """ tokenizer = AutoTokenizer.from_pretrained(model_name) prompt = "The meaning of life is" # Reference: unpatched model ref_model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.bfloat16, device_map="cuda", ) inputs = tokenizer(prompt, return_tensors="pt").to("cuda") with torch.no_grad(): ref_logits = ref_model(**inputs).logits # Test: patched model test_model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.bfloat16, device_map="cuda", ) inject_fn(test_model, cuda_rmsnorm_fn) with torch.no_grad(): test_logits = test_model(**inputs).logits # Compare logits max_diff = (ref_logits - test_logits).abs().max().item() mean_diff = (ref_logits - test_logits).abs().mean().item() print(f"Logits comparison:") print(f" Max absolute difference: {max_diff:.8f}") print(f" Mean absolute difference: {mean_diff:.8f}") # Check that argmax (predicted tokens) match ref_tokens = ref_logits.argmax(dim=-1) test_tokens = test_logits.argmax(dim=-1) token_match = (ref_tokens == test_tokens).all().item() print(f" Token predictions match: {token_match}") # Also compare full generation ref_gen = ref_model.generate(**inputs, max_new_tokens=max_new_tokens, do_sample=False) test_gen = test_model.generate(**inputs, max_new_tokens=max_new_tokens, do_sample=False) ref_text = tokenizer.decode(ref_gen[0], skip_special_tokens=True) test_text = tokenizer.decode(test_gen[0], skip_special_tokens=True) print(f"\nReference output: {ref_text}") print(f"Test output: {test_text}") print(f"Outputs match: {ref_text == test_text}") return max_diff < 1e-2 and token_match # Usage verify_transformers_correctness( "meta-llama/Llama-2-7b-hf", inject_transformers_kernels, cuda_rmsnorm, ) ``` ### Module-Level Verification ```python def verify_rmsnorm_module(model_name="meta-llama/Llama-2-7b-hf"): """Verify a single RMSNorm module against the reference.""" from transformers.models.llama.modeling_llama import LlamaRMSNorm # Create reference module config = AutoConfig.from_pretrained(model_name) hidden_size = config.hidden_size eps = config.rms_norm_eps ref_norm = LlamaRMSNorm(hidden_size, eps=eps).cuda().to(torch.bfloat16) # Test with random inputs x = torch.randn(2, 128, hidden_size, dtype=torch.bfloat16, device="cuda") with torch.no_grad(): ref_out = ref_norm(x) custom_out = cuda_rmsnorm(x, ref_norm.weight, eps) max_diff = (ref_out - custom_out).abs().max().item() print(f"LlamaRMSNorm max diff: {max_diff:.8f}") assert max_diff < 1e-2, f"Verification failed: max_diff={max_diff}" print("Verification PASSED") ``` ### Performance Verification ```python import time def benchmark_transformers_model( model_name: str, inject_fn=None, prompt: str = "Explain quantum computing in simple terms:", max_new_tokens: int = 100, num_runs: int = 5, ): """Benchmark a transformers model with and without custom kernels.""" tokenizer = AutoTokenizer.from_pretrained(model_name) tokenizer.pad_token = tokenizer.eos_token model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.bfloat16, device_map="cuda", ) inputs = tokenizer(prompt, return_tensors="pt").to("cuda") if inject_fn: inject_fn(model, cuda_rmsnorm) # Warmup with torch.no_grad(): model.generate(**inputs, max_new_tokens=10, do_sample=False) # Benchmark times = [] with torch.no_grad(): for _ in range(num_runs): torch.cuda.synchronize() start = time.perf_counter() model.generate(**inputs, max_new_tokens=max_new_tokens, do_sample=False) torch.cuda.synchronize() end = time.perf_counter() times.append(end - start) avg_time = sum(times) / len(times) tokens_per_sec = max_new_tokens / avg_time print(f"Average time: {avg_time:.3f}s") print(f"Tokens/sec: {tokens_per_sec:.1f}") return avg_time, tokens_per_sec ``` ## Complete Integration Example ```python import torch from transformers import AutoModelForCausalLM, AutoTokenizer from huggingface_kernels import get_kernel def main(): model_name = "meta-llama/Llama-2-7b-hf" # Step 1: Load model print("Loading model...") model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.bfloat16, device_map="cuda", ) tokenizer = AutoTokenizer.from_pretrained(model_name) # Step 2: Load custom kernel print("Loading custom kernels...") rmsnorm_kernel = get_kernel("huggingface/cuda-kernels", "rmsnorm") def cuda_rmsnorm(x, weight, eps): return rmsnorm_kernel.rmsnorm_forward(x, weight, eps) # Step 3: Inject kernels print("Injecting kernels...") inject_transformers_kernels(model, cuda_rmsnorm) # Step 4: Optionally compile # model = torch.compile(model, mode="reduce-overhead") # Step 5: Run inference prompt = "The future of artificial intelligence is" inputs = tokenizer(prompt, return_tensors="pt").to("cuda") with torch.no_grad(): output = model.generate( **inputs, max_new_tokens=100, do_sample=False, ) text = tokenizer.decode(output[0], skip_special_tokens=True) print(f"\nGenerated text:\n{text}") if __name__ == "__main__": main() ``` ## Summary When integrating with transformers: 1. **Use `variance_epsilon`** not `eps` for the epsilon attribute 2. **Weight is always present** -- no need for None checks (unlike diffusers) 3. **No `set_processor`** -- patch forward methods directly for attention 4. **Handle `device_map`** -- modules may be on different devices 5. **Import model-specific classes** or use the universal detection approach 6. **Test with `generate()`** to verify end-to-end correctness through autoregressive decoding 7. **Flash Attention 2 works alongside** custom kernels for other operations