test / skill_example /references /transformers-integration.md
Jack-Khuu
Demo
88a1dd2
|
Raw History Blame Contribute Delete
27.6 kB

A newer version of the Gradio SDK is available: 6.29.1

Upgrade

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:

  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