Spaces:
Sleeping
Download skill_example/references/transformers-integration.md from KhookieThief/test: direct link, hf CLI and curl.
- Browser
- Download file 27.6 kB
-
https://huggingface.co/spaces/KhookieThief/test/resolve/main/skill_example/references/transformers-integration.md
- Command line
-
hf download hf://spaces/KhookieThief/test/skill_example/references/transformers-integration.md
-
curl -L -o transformers-integration.md https://huggingface.co/spaces/KhookieThief/test/resolve/main/skill_example/references/transformers-integration.md
A newer version of the Gradio SDK is available: 6.29.1
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
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
# 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:
# 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
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:
# 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:
# 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:
# 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:
# 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
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
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
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
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
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
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
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
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
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
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:
- Use
variance_epsilonnotepsfor the epsilon attribute - Weight is always present -- no need for None checks (unlike diffusers)
- No
set_processor-- patch forward methods directly for attention - Handle
device_map-- modules may be on different devices - Import model-specific classes or use the universal detection approach
- Test with
generate()to verify end-to-end correctness through autoregressive decoding - Flash Attention 2 works alongside custom kernels for other operations