This commit is contained in:
Daniel Han 2026-02-05 06:10:06 -08:00
commit 5000413815
14 changed files with 1298 additions and 95 deletions

View file

@ -0,0 +1,500 @@
# Unsloth
# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Affero General Public License as published
# by the Free Software Foundation, either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# GNU Affero General Public License for more details.
#
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.
"""
Auto-tuning cache system for MoE kernels to ensure tuning runs only once at training start.
"""
import hashlib
import json
import logging
import os
import time
from typing import Dict, List, Optional, Tuple, Any
import torch
import triton
logger = logging.getLogger(__name__)
# Global cache for kernel configurations
_kernel_config_cache: Dict[str, Any] = {}
_autotune_completed: Dict[str, bool] = {}
def _get_cache_key(
num_experts: int,
hidden_dim: int,
intermediate_dim: int,
top_k: int,
dtype: torch.dtype,
device_capability: Tuple[int, int],
seq_len: int = 8192, # Default sequence length for tuning
) -> str:
"""Generate a unique cache key based on model configuration."""
key_data = {
"num_experts": num_experts,
"hidden_dim": hidden_dim,
"intermediate_dim": intermediate_dim,
"top_k": top_k,
"dtype": str(dtype),
"device_capability": device_capability,
"seq_len": seq_len,
}
key_str = json.dumps(key_data, sort_keys = True)
return hashlib.md5(key_str.encode()).hexdigest()
def _get_cache_file_path(cache_key: str) -> str:
"""Get the file path for the cache file."""
cache_dir = os.path.expanduser("~/.cache/unsloth/moe_autotune")
os.makedirs(cache_dir, exist_ok = True)
return os.path.join(cache_dir, f"{cache_key}.json")
def load_cached_config(cache_key: str) -> Optional[Dict[str, Any]]:
"""Load cached kernel configuration from disk."""
cache_file = _get_cache_file_path(cache_key)
if not os.path.exists(cache_file):
return None
try:
with open(cache_file, "r") as f:
cached_data = json.load(f)
# Verify cache is still valid (same device, etc.)
current_device_capability = torch.cuda.get_device_capability()
if cached_data.get("device_capability") != current_device_capability:
logger.info("Device capability changed, invalidating cache")
os.remove(cache_file)
return None
logger.info(f"Loaded cached MoE kernel config: {cache_key}")
return cached_data
except Exception as e:
logger.warning(f"Failed to load cache file {cache_file}: {e}")
try:
os.remove(cache_file)
except:
pass
return None
def save_cached_config(
cache_key: str,
config_fwd: Any,
config_bwd_dx: Any,
config_bwd_dw: Any,
metadata: Dict[str, Any] = None,
) -> None:
"""Save kernel configuration to disk cache."""
cache_file = _get_cache_file_path(cache_key)
cache_data = {
"timestamp": time.time(),
"device_capability": torch.cuda.get_device_capability(),
"config_fwd": config_fwd.__dict__
if hasattr(config_fwd, "__dict__")
else str(config_fwd),
"config_bwd_dx": config_bwd_dx.__dict__
if hasattr(config_bwd_dx, "__dict__")
else str(config_bwd_dx),
"config_bwd_dw": config_bwd_dw.__dict__
if hasattr(config_bwd_dw, "__dict__")
else str(config_bwd_dw),
"metadata": metadata or {},
}
try:
with open(cache_file, "w") as f:
json.dump(cache_data, f, indent = 2)
logger.info(f"Saved MoE kernel config cache: {cache_key}")
except Exception as e:
logger.warning(f"Failed to save cache file {cache_file}: {e}")
def get_or_autotune_moe_kernels(
num_experts: int,
hidden_dim: int,
intermediate_dim: int,
top_k: int,
dtype: torch.dtype,
force_autotune: bool = False,
seq_len: int = 8192,
) -> Tuple[Any, Any, Any]:
"""
Get cached kernel configurations or run auto-tuning.
Args:
num_experts: Number of experts in the MoE layer
hidden_dim: Hidden dimension of the model
intermediate_dim: Intermediate dimension for MoE MLP
top_k: Number of experts to route to
dtype: Data type for computation
force_autotune: Force re-running autotuning even if cache exists
seq_len: Sequence length to use for tuning benchmarks
Returns:
Tuple of (config_fwd, config_bwd_dx, config_bwd_dw)
"""
device_capability = torch.cuda.get_device_capability()
cache_key = _get_cache_key(
num_experts,
hidden_dim,
intermediate_dim,
top_k,
dtype,
device_capability,
seq_len,
)
# 0. Check for environment variable override to DISABLE autotuning
if os.environ.get("UNSLOTH_MOE_DISABLE_AUTOTUNE", "0") == "1":
logger.info(
f"UNSLOTH_MOE_DISABLE_AUTOTUNE=1: Using Heuristic (Safe) MoE kernel configs for SM{device_capability[0]}{device_capability[1]}"
)
return _get_heuristic_configs()
if not force_autotune and cache_key in _kernel_config_cache:
logger.info(f"Using in-memory cached MoE kernel configs: {cache_key}")
return _kernel_config_cache[cache_key]
# Try to load from disk
if not force_autotune:
cached_data = load_cached_config(cache_key)
if cached_data is not None:
# Reconstruct config objects from cached data
try:
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
)
config_fwd = KernelConfigForward(**cached_data["config_fwd"])
config_bwd_dx = KernelConfigBackward_dX(**cached_data["config_bwd_dx"])
config_bwd_dw = KernelConfigBackward_dW(**cached_data["config_bwd_dw"])
configs = (config_fwd, config_bwd_dx, config_bwd_dw)
_kernel_config_cache[cache_key] = configs
return configs
except Exception as e:
logger.warning(f"Failed to reconstruct cached configs: {e}")
# Run autotuning
if cache_key in _autotune_completed and not force_autotune:
logger.info(f"Autotuning already completed for: {cache_key}")
return _kernel_config_cache[cache_key]
logger.info(f"Running MoE kernel auto-tuning for: {cache_key}")
logger.info(
f"Configuration: {num_experts} experts, {hidden_dim} hidden, {intermediate_dim} intermediate, top_k={top_k}"
)
try:
configs = _run_moe_autotuning(
num_experts, hidden_dim, intermediate_dim, top_k, dtype, seq_len
)
# Cache the results
_kernel_config_cache[cache_key] = configs
_autotune_completed[cache_key] = True
# Save to disk
config_fwd, config_bwd_dx, config_bwd_dw = configs
save_cached_config(
cache_key,
config_fwd,
config_bwd_dx,
config_bwd_dw,
{
"num_experts": num_experts,
"hidden_dim": hidden_dim,
"intermediate_dim": intermediate_dim,
},
)
logger.info(f"MoE kernel auto-tuning completed: {cache_key}")
return configs
except Exception as e:
logger.error(f"MoE kernel auto-tuning failed: {e}")
if "AttributeError" in str(e) and "_experimental_make_tensor_descriptor" in str(
e
):
logger.warning(
"Unsloth: Your Triton version might be incompatible with TMA features. Falling back to default configs."
)
logger.info("Falling back to default kernel configurations")
return _get_default_configs()
def _run_moe_autotuning(
num_experts: int,
hidden_dim: int,
intermediate_dim: int,
top_k: int,
dtype: torch.dtype,
seq_len: int,
) -> Tuple[Any, Any, Any]:
"""Run the actual auto-tuning for MoE kernels."""
# Create dummy inputs for tuning
device = "cuda"
# Use a fixed, safe number of tokens for autotuning to avoid OOMs and dependency on seq_len
# 4096 is standard for finding good kernels without consuming 10GB+ VRAM
# We ignore the passed seq_len for the actual allocation to satisfy user request
num_tokens = 4096
total_tokens = num_tokens * top_k
# Create dummy tensors
hidden_states = torch.randn(num_tokens, hidden_dim, device = device, dtype = dtype)
# Create dummy weights
gate_up_weights = torch.randn(
num_experts, 2 * intermediate_dim, hidden_dim, device = device, dtype = dtype
)
down_weights = torch.randn(
num_experts, hidden_dim, intermediate_dim, device = device, dtype = dtype
)
# Create dummy routing data
m_sizes = torch.randint(
1, total_tokens // num_experts + 1, (num_experts,), device = device
)
m_sizes = m_sizes * (total_tokens // m_sizes.sum().item())
# Adjust to ensure exact total
diff = total_tokens - m_sizes.sum().item()
if diff != 0:
m_sizes[0] += diff
gather_indices = torch.arange(total_tokens, device = device)
torch.randperm(total_tokens, out = gather_indices)
# Autotune forward kernel - use the interface function with autotune=True
# This properly invokes the kernel and lets triton handle the autotuning
from .grouped_gemm.interface import (
grouped_gemm_forward,
grouped_gemm_dX,
grouped_gemm_dW,
)
from .grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
from .grouped_gemm.kernels.backward import (
_autotuned_grouped_gemm_dX_kernel,
_autotuned_grouped_gemm_dW_kernel,
)
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
)
logger.info("Autotuning forward kernel (first GEMM)...")
# Run with autotune=True to trigger autotuning
_ = grouped_gemm_forward(
X = hidden_states,
W = gate_up_weights,
topk = top_k,
m_sizes = m_sizes,
gather_indices = gather_indices,
permute_x = True,
permute_y = False,
autotune = True,
)
triton_config_fwd = _autotuned_grouped_gemm_forward_kernel.best_config
# Convert triton.Config to KernelConfigForward
config_fwd = KernelConfigForward(
BLOCK_SIZE_M = triton_config_fwd.kwargs["BLOCK_SIZE_M"],
BLOCK_SIZE_N = triton_config_fwd.kwargs["BLOCK_SIZE_N"],
BLOCK_SIZE_K = triton_config_fwd.kwargs["BLOCK_SIZE_K"],
num_warps = triton_config_fwd.num_warps,
num_stages = triton_config_fwd.num_stages,
use_tma_load_x = triton_config_fwd.kwargs.get("USE_TMA_LOAD_X", False),
use_tma_load_w = triton_config_fwd.kwargs.get("USE_TMA_LOAD_W", False),
use_tma_store = triton_config_fwd.kwargs.get("USE_TMA_STORE", False),
)
# Autotune backward dX kernel
logger.info("Autotuning backward dX kernel...")
dummy_grad = torch.randn(
total_tokens, 2 * intermediate_dim, device = device, dtype = dtype
)
_ = grouped_gemm_dX(
dY = dummy_grad,
W = gate_up_weights,
gather_indices = gather_indices,
m_sizes = m_sizes,
topk = top_k,
permute_x = True,
permute_y = False,
autotune = True,
)
triton_config_bwd_dx = _autotuned_grouped_gemm_dX_kernel.best_config
# Convert triton.Config to KernelConfigBackward_dX
config_bwd_dx = KernelConfigBackward_dX(
BLOCK_SIZE_M = triton_config_bwd_dx.kwargs["BLOCK_SIZE_M"],
BLOCK_SIZE_N = triton_config_bwd_dx.kwargs["BLOCK_SIZE_N"],
BLOCK_SIZE_K = triton_config_bwd_dx.kwargs["BLOCK_SIZE_K"],
num_warps = triton_config_bwd_dx.num_warps,
num_stages = triton_config_bwd_dx.num_stages,
use_tma_load_dy = triton_config_bwd_dx.kwargs.get("USE_TMA_LOAD_dY", False),
use_tma_load_w = triton_config_bwd_dx.kwargs.get("USE_TMA_LOAD_W", False),
use_tma_store = triton_config_bwd_dx.kwargs.get("USE_TMA_STORE", False),
)
# Autotune backward dW kernel
logger.info("Autotuning backward dW kernel...")
_ = grouped_gemm_dW(
X = hidden_states,
dY = dummy_grad,
m_sizes = m_sizes,
gather_indices = gather_indices,
topk = top_k,
permute_x = True,
permute_y = False,
autotune = True,
)
triton_config_bwd_dw = _autotuned_grouped_gemm_dW_kernel.best_config
# Convert triton.Config to KernelConfigBackward_dW
config_bwd_dw = KernelConfigBackward_dW(
BLOCK_SIZE_M = triton_config_bwd_dw.kwargs["BLOCK_SIZE_M"],
BLOCK_SIZE_N = triton_config_bwd_dw.kwargs["BLOCK_SIZE_N"],
BLOCK_SIZE_K = triton_config_bwd_dw.kwargs["BLOCK_SIZE_K"],
num_warps = triton_config_bwd_dw.num_warps,
num_stages = triton_config_bwd_dw.num_stages,
use_tma_load_dy = triton_config_bwd_dw.kwargs.get("USE_TMA_LOAD_dY", False),
use_tma_load_x = triton_config_bwd_dw.kwargs.get("USE_TMA_LOAD_X", False),
use_tma_store = triton_config_bwd_dw.kwargs.get("USE_TMA_STORE", False),
)
return config_fwd, config_bwd_dx, config_bwd_dw
return config_fwd, config_bwd_dx, config_bwd_dw
def _get_heuristic_configs() -> Tuple[Any, Any, Any]:
"""
Get 'Safe Heuristic' kernel configurations.
These are verified to be safe on A100 (SM80) and provide ~9x speedup on H100/B200.
"""
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
)
# Safe Forward Config: 64x128x128 (Fits A100 SMEM)
config_fwd = KernelConfigForward(
BLOCK_SIZE_M = 64,
BLOCK_SIZE_N = 128,
BLOCK_SIZE_K = 128,
num_warps = 8,
num_stages = 3,
permute_x = True,
permute_y = True,
use_tma_load_x = False,
use_tma_load_w = False, # TMA loads might need alignment checks, safer to disable for heuristic
use_tma_store = False,
)
# Safe Backward Configs: 64x64x256
config_bwd_dx = KernelConfigBackward_dX(
BLOCK_SIZE_M = 64,
BLOCK_SIZE_N = 64,
BLOCK_SIZE_K = 256,
num_warps = 8,
num_stages = 4,
permute_x = True,
permute_y = True,
use_tma_load_dy = False,
use_tma_load_w = False,
use_tma_store = False,
)
config_bwd_dw = KernelConfigBackward_dW(
BLOCK_SIZE_M = 64,
BLOCK_SIZE_N = 64,
BLOCK_SIZE_K = 256,
num_warps = 8,
num_stages = 4,
permute_x = True,
permute_y = True,
use_tma_load_dy = False,
use_tma_load_x = False,
use_tma_store = False,
)
return config_fwd, config_bwd_dx, config_bwd_dw
def _get_default_configs() -> Tuple[Any, Any, Any]:
"""Get default kernel configurations as fallback."""
from .grouped_gemm.kernels.tuning import (
KernelConfigForward,
KernelConfigBackward_dX,
KernelConfigBackward_dW,
)
logger.warning("Using default MoE kernel configurations (not optimal)")
config_fwd = KernelConfigForward(
BLOCK_SIZE_M = 128,
BLOCK_SIZE_N = 128,
BLOCK_SIZE_K = 64,
num_warps = 8,
num_stages = 3,
use_tma_load_x = False,
use_tma_load_w = False,
use_tma_store = False,
)
config_bwd_dx = KernelConfigBackward_dX(
BLOCK_SIZE_M = 128,
BLOCK_SIZE_N = 128,
BLOCK_SIZE_K = 64,
num_warps = 8,
num_stages = 3,
use_tma_load_dy = False,
use_tma_load_w = False,
use_tma_store = False,
)
config_bwd_dw = KernelConfigBackward_dW(
BLOCK_SIZE_M = 128,
BLOCK_SIZE_N = 128,
BLOCK_SIZE_K = 64,
num_warps = 8,
num_stages = 3,
use_tma_load_dy = False,
use_tma_load_x = False,
use_tma_store = False,
)
return config_fwd, config_bwd_dx, config_bwd_dw
def clear_cache() -> None:
"""Clear all cached kernel configurations."""
global _kernel_config_cache, _autotune_completed
_kernel_config_cache.clear()
_autotune_completed.clear()
logger.info("Cleared MoE kernel cache")
def is_autotuning_completed(cache_key: str) -> bool:
"""Check if autotuning has been completed for a given cache key."""
return cache_key in _autotune_completed

View file

@ -8,17 +8,17 @@ from dataclasses import asdict
import torch
import triton
from grouped_gemm.kernels.backward import (
from .kernels.backward import (
_autotuned_grouped_gemm_dW_kernel,
_autotuned_grouped_gemm_dX_kernel,
_grouped_gemm_dW_kernel,
_grouped_gemm_dX_kernel,
)
from grouped_gemm.kernels.forward import (
from .kernels.forward import (
_autotuned_grouped_gemm_forward_kernel,
_grouped_gemm_forward_kernel,
)
from grouped_gemm.kernels.tuning import (
from .kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
@ -35,17 +35,57 @@ ch = logging.StreamHandler()
ch.setFormatter(formatter)
logger.addHandler(ch)
_FUSED_MUL_WARN = False
_SUPPORTS_TMA = None
# Precompute TMA support to avoid graph breaks
# TMA requires both:
# 1. GPU capability >= 9 (Hopper+)
# 2. Triton version with TMA API (make_tensor_descriptor or _experimental_make_tensor_descriptor)
def _check_tma_support():
import triton.language as tl
gpu_supports_tma = torch.cuda.get_device_capability()[0] >= 9
# Check for both old experimental and new stable API names
triton_has_tma_api = hasattr(tl, "make_tensor_descriptor") or hasattr(
tl, "_experimental_make_tensor_descriptor"
)
return gpu_supports_tma and triton_has_tma_api
_SUPPORTS_TMA = _check_tma_support()
# Check if triton.set_allocator is available (Triton 3.0+)
_HAS_SET_ALLOCATOR = hasattr(triton, "set_allocator")
def supports_tma():
global _SUPPORTS_TMA
if _SUPPORTS_TMA is None:
_SUPPORTS_TMA = torch.cuda.get_device_capability()[0] >= 9
return _SUPPORTS_TMA
# Helper to support allow_in_graph
try:
from torch.compiler import allow_in_graph
except ImportError:
from torch._dynamo import allow_in_graph
# Helper to detect if we're in tracing/compilation mode
def _is_tracing(*tensors):
"""
Check if tensors are fake tensors used during torch.compile tracing.
During tracing, tensors are FakeTensor/FunctionalTensor and we can't run Triton kernels.
During execution, tensors are real Tensors and we MUST run the kernels.
NOTE: We do NOT use torch.compiler.is_compiling() because it returns True
during both tracing AND execution. We only want to skip kernels during tracing
when tensors are actually fake.
"""
for t in tensors:
name = type(t).__name__
if name in ("FakeTensor", "FunctionalTensor", "FunctionalTensorWrapper"):
return True
return False
_per_device_alloc_fns = {}
@ -83,6 +123,7 @@ def log_kernel_info(
logger.debug(f"{kernel_name} autotuned best_config: {best_config}")
@allow_in_graph
def grouped_gemm_forward(
X: torch.Tensor,
W: torch.Tensor,
@ -158,11 +199,21 @@ def grouped_gemm_forward(
use_tma_store = False
if use_tma or autotune:
# Respect global persistent allocator if set
if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False):
def alloc_fn(size: int, alignment: int, stream: int):
return torch.empty(size, device = "cuda", dtype = torch.int8)
def alloc_fn(size: int, alignment: int, stream: int):
return torch.empty(size, device = "cuda", dtype = torch.int8)
triton.set_allocator(alloc_fn)
triton.set_allocator(alloc_fn)
if W.ndim == 3:
num_experts = W.shape[0]
N = W.shape[1]
# K = W.shape[2]
else:
num_experts = m_sizes.shape[0]
N = W.shape[0] // num_experts
X = X.view(-1, X.shape[-1])
W = W.view(-1, W.shape[-1])
@ -188,9 +239,7 @@ def grouped_gemm_forward(
total_tokens = X.shape[0]
num_tokens = total_tokens // topk
num_experts = m_sizes.shape[0]
_, K = X.shape
N = W.shape[0] // num_experts
assert K == W.shape[1], f"K ({K}) must match W.shape[1] ({W.shape[1]})"
if fuse_mul_post:
@ -212,8 +261,8 @@ def grouped_gemm_forward(
)
y = torch.empty((total_tokens, N), device = X.device, dtype = X.dtype)
if total_tokens == 0 or N == 0:
return y
# if total_tokens == 0 or N == 0:
# return y
NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count
@ -221,9 +270,9 @@ def grouped_gemm_forward(
return (NUM_SMS,)
if not autotune:
BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
BLOCK_SIZE_N = min(N, BLOCK_SIZE_N)
BLOCK_SIZE_M = min(total_tokens, BLOCK_SIZE_M)
# BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
# BLOCK_SIZE_N = min(N, BLOCK_SIZE_N)
pass
if debug:
print(
@ -276,16 +325,19 @@ def grouped_gemm_forward(
if autotune
else _grouped_gemm_forward_kernel
)
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
is_fake = _is_tracing(X, W)
if not is_fake:
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
return y
@allow_in_graph
def grouped_gemm_dX(
dY: torch.Tensor,
W: torch.Tensor,
@ -354,20 +406,28 @@ def grouped_gemm_dX(
use_tma_store = False
if use_tma or autotune:
# Respect global persistent allocator if set
if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False):
def alloc_fn(size: int, alignment: int, stream: int):
# print(f"DEBUG::GROUPED_GEMM alloc_fn {size=} {alignment=} {stream=}")
return torch.empty(size, device = "cuda", dtype = torch.int8)
def alloc_fn(size: int, alignment: int, stream: int):
# print(f"DEBUG::GROUPED_GEMM alloc_fn {size=} {alignment=} {stream=}")
return torch.empty(size, device = "cuda", dtype = torch.int8)
triton.set_allocator(alloc_fn)
triton.set_allocator(alloc_fn)
if W.ndim == 3:
num_experts = W.shape[0]
N = W.shape[1]
else:
num_experts = m_sizes.shape[0]
N = W.shape[0] // num_experts
num_experts = m_sizes.shape[0]
dY = dY.view(-1, dY.shape[-1])
W = W.view(-1, W.shape[-1])
M_total, N_grad = dY.shape
N_total, K = W.shape
N = N_total // num_experts
# N = N_total // num_experts
assert N_grad == N, f"Grad_output N ({N_grad}) must match weight N ({N})"
assert (
@ -393,9 +453,9 @@ def grouped_gemm_dX(
return (NUM_SMS,)
if not autotune:
BLOCK_SIZE_M = min(M_total, BLOCK_SIZE_M)
BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N)
BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
# BLOCK_SIZE_N = min(N_grad, BLOCK_SIZE_N)
# BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
pass
if debug:
print(
@ -437,15 +497,19 @@ def grouped_gemm_dX(
}
)
kernel = _autotuned_grouped_gemm_dX_kernel if autotune else _grouped_gemm_dX_kernel
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
is_fake = _is_tracing(dY, W)
if not is_fake:
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
return dX
@allow_in_graph
def grouped_gemm_dW(
X: torch.Tensor,
dY: torch.Tensor,
@ -510,11 +574,13 @@ def grouped_gemm_dW(
use_tma_store = False
if use_tma or autotune:
# Respect global persistent allocator if set
if _HAS_SET_ALLOCATOR and not getattr(triton, "_unsloth_allocator_set", False):
def alloc_fn(size: int, alignment: int, stream: int):
return torch.empty(size, device = "cuda", dtype = torch.int8)
def alloc_fn(size: int, alignment: int, stream: int):
return torch.empty(size, device = "cuda", dtype = torch.int8)
triton.set_allocator(alloc_fn)
triton.set_allocator(alloc_fn)
if permute_x or permute_y:
assert gather_indices is not None
@ -541,9 +607,9 @@ def grouped_gemm_dW(
dW = torch.zeros((num_experts, N, K), device = X.device, dtype = X.dtype)
if not autotune:
BLOCK_SIZE_M = min(total_tokens, BLOCK_SIZE_M)
BLOCK_SIZE_N = min(N, BLOCK_SIZE_N)
BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
# BLOCK_SIZE_N = min(N, BLOCK_SIZE_N)
# BLOCK_SIZE_K = min(K, BLOCK_SIZE_K)
pass
def grid(META):
return (NUM_SMS,)
@ -607,12 +673,15 @@ def grouped_gemm_dW(
)
kernel = _autotuned_grouped_gemm_dW_kernel if autotune else _grouped_gemm_dW_kernel
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
is_fake = _is_tracing(X, dY)
if not is_fake:
compiled_kernel: triton.compiler.CompiledKernel = kernel[grid](**kernel_args)
if autotune:
log_kernel_info(compiled_kernel, kernel.best_config)
else:
log_kernel_info(compiled_kernel)
return dW
@ -680,6 +749,7 @@ class GroupedGemm(torch.autograd.Function):
@staticmethod
def backward(ctx, dY):
dY = dY.contiguous()
X, W, m_sizes, gather_indices = ctx.saved_tensors
topk = ctx.topk
permute_x = ctx.permute_x

View file

@ -1,5 +1,18 @@
# SPDX-License-Identifier: GNU Affero General Public License v3.0
# Copyright 2023-present the Unsloth team. All rights reserved.
# Unsloth
# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Affero General Public License as published
# by the Free Software Foundation, either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# GNU Affero General Public License for more details.
#
# You should have received a copy of the GNU Affero General Public License
# along with this program. If not, see <https://www.gnu.org/licenses/>.
"""
Autotuning utils
@ -36,17 +49,39 @@ def convert_args_to_list(args):
return [val_to_list(arg) for arg in args]
def _triton_supports_tma():
"""Check if current Triton version supports TMA API."""
import triton.language as tl
# Check for both old experimental and new stable API names
return hasattr(tl, "make_tensor_descriptor") or hasattr(
tl, "_experimental_make_tensor_descriptor"
)
# Precompute at module import
# NOTE: TMA is disabled for now due to compatibility issues with permute_x/permute_y settings
# in the MoE grouped GEMM forward/backward passes. Re-enable once these are resolved.
_TRITON_HAS_TMA = False # _triton_supports_tma()
def get_forward_configs(
BLOCK_M = DEFAULT_M_BLOCK_SIZES,
BLOCK_N = DEFAULT_N_BLOCK_SIZES,
BLOCK_K = DEFAULT_K_BLOCK_SIZES,
TMA_LOAD_X = True,
TMA_LOAD_W = True,
TMA_LOAD_X = None, # Auto-detect if not specified
TMA_LOAD_W = None, # Auto-detect if not specified
TMA_STORE = False, # NOTE: TMA_STORE is disabled for now
num_warps = DEFAULT_NUM_WARPS,
num_stages = DEFAULT_NUM_STAGES,
num_ctas = DEFAULT_NUM_CTAS,
):
# Auto-detect TMA support
if TMA_LOAD_X is None:
TMA_LOAD_X = _TRITON_HAS_TMA
if TMA_LOAD_W is None:
TMA_LOAD_W = _TRITON_HAS_TMA
(
BLOCK_M,
BLOCK_N,
@ -115,13 +150,18 @@ def get_dX_kernel_configs(
BLOCK_M = DEFAULT_M_BLOCK_SIZES,
BLOCK_N = DEFAULT_N_BLOCK_SIZES,
BLOCK_K = DEFAULT_K_BLOCK_SIZES,
TMA_LOAD_dY = True,
TMA_LOAD_W = True,
TMA_LOAD_dY = None, # Auto-detect if not specified
TMA_LOAD_W = None, # Auto-detect if not specified
TMA_STORE = False, # NOTE: TMA_STORE is disabled for now
num_warps = DEFAULT_NUM_WARPS,
num_stages = DEFAULT_NUM_STAGES,
num_ctas = DEFAULT_NUM_CTAS,
):
# Auto-detect TMA support
if TMA_LOAD_dY is None:
TMA_LOAD_dY = _TRITON_HAS_TMA
if TMA_LOAD_W is None:
TMA_LOAD_W = _TRITON_HAS_TMA
(
BLOCK_M,
BLOCK_N,
@ -193,10 +233,15 @@ def get_dW_kernel_configs(
num_warps = DEFAULT_NUM_WARPS,
num_stages = DEFAULT_NUM_STAGES,
num_ctas = DEFAULT_NUM_CTAS,
TMA_LOAD_dY = True,
TMA_LOAD_X = True,
TMA_LOAD_dY = None, # Auto-detect if not specified
TMA_LOAD_X = None, # Auto-detect if not specified
TMA_STORE = False,
):
# Auto-detect TMA support
if TMA_LOAD_dY is None:
TMA_LOAD_dY = _TRITON_HAS_TMA
if TMA_LOAD_X is None:
TMA_LOAD_X = _TRITON_HAS_TMA
(
BLOCK_M,
BLOCK_N,
@ -291,8 +336,8 @@ def exceeds_smem_capacity(
def common_prune_criteria(config: triton.Config, kwargs: dict, dtype):
from grouped_gemm.interface import supports_tma
from grouped_gemm.kernels.tuning import get_device_properties
from ..interface import supports_tma
from .tuning import get_device_properties
smem_size = get_device_properties().SIZE_SMEM
@ -323,7 +368,7 @@ def common_prune_criteria(config: triton.Config, kwargs: dict, dtype):
def maybe_disable_tma(config: triton.Config):
from grouped_gemm.interface import supports_tma
from ..interface import supports_tma
tma_keys = [k for k in config.kwargs.keys() if k.startswith("USE_TMA_")]
if not supports_tma():

View file

@ -5,7 +5,7 @@ import torch
import triton
import triton.language as tl
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
get_dW_kernel_configs,
get_dX_kernel_configs,
prune_dX_configs,
@ -53,11 +53,11 @@ def _grouped_gemm_dX_kernel(
m_sizes_ptr,
# problem sizes
NUM_EXPERTS: tl.constexpr,
NUM_TOKENS: tl.constexpr,
NUM_TOKENS,
TOPK: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
NUM_SMS: tl.constexpr,
NUM_SMS,
# Tuning parameters
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
@ -69,7 +69,7 @@ def _grouped_gemm_dX_kernel(
USE_TMA_STORE: tl.constexpr = False,
FLATTEN: tl.constexpr = True,
) -> None:
TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK
TOTAL_TOKENS = NUM_TOKENS * TOPK
output_dtype = dX_ptr.dtype.element_ty
tidx = tl.program_id(0)
@ -82,7 +82,7 @@ def _grouped_gemm_dX_kernel(
# Also, we are defining a single global descriptor with single block shape
# Need to check that this does not result in errors when crossing expert boundaries
if USE_TMA_LOAD_dY:
dY_desc = tl._experimental_make_tensor_descriptor(
dY_desc = tl.make_tensor_descriptor(
dY_ptr,
shape = [TOTAL_TOKENS, N],
strides = [N, 1],
@ -91,7 +91,7 @@ def _grouped_gemm_dX_kernel(
if USE_TMA_LOAD_W:
expert_stride = N * K
w_desc = tl._experimental_make_tensor_descriptor(
w_desc = tl.make_tensor_descriptor(
w_ptr,
shape = [NUM_EXPERTS, N, K],
strides = [expert_stride, K, 1],
@ -123,7 +123,7 @@ def _grouped_gemm_dX_kernel(
tl.static_assert(
K % BLOCK_SIZE_K == 0, "K must be divisible by BLOCK_SIZE_K"
)
dX_desc = tl._experimental_make_tensor_descriptor(
dX_desc = tl.make_tensor_descriptor(
dX_ptr,
shape = [m_end, K],
strides = [K, 1],
@ -232,6 +232,7 @@ def _grouped_gemm_dX_kernel(
# TODO: check if predication along K is needed since we checked that K is divisible by BLOCK_SIZE_K in the forward kernel
# [M, N] @ [N, K] -> [M, K]
dY = dY.to(w.dtype)
accumulator += tl.dot(dY, w) # NOTE: no transpose of b
# Advance A along contiguous dimension
@ -266,7 +267,8 @@ def _grouped_gemm_dX_kernel(
_autotuned_grouped_gemm_dX_kernel = triton.autotune(
configs = get_dX_kernel_configs(),
prune_configs_by = {"early_config_prune": prune_dX_configs},
key = ["NUM_EXPERTS", "NUM_TOKENS", "N", "K", "PERMUTE_X", "PERMUTE_Y"],
# NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length
key = ["NUM_EXPERTS", "N", "K", "PERMUTE_X", "PERMUTE_Y"],
)(_grouped_gemm_dX_kernel)
"""
@ -298,12 +300,12 @@ def _grouped_gemm_dW_kernel(
m_sizes_ptr,
gather_indices_ptr,
# problem sizes
NUM_TOKENS: tl.constexpr,
NUM_TOKENS,
TOPK: tl.constexpr,
NUM_EXPERTS: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
NUM_SMS: tl.constexpr,
NUM_SMS,
BLOCK_SIZE_N: tl.constexpr,
BLOCK_SIZE_K: tl.constexpr,
BLOCK_SIZE_M: tl.constexpr,
@ -315,14 +317,14 @@ def _grouped_gemm_dW_kernel(
FLATTEN: tl.constexpr = True,
acc_dtype: tl.constexpr = tl.float32,
) -> None:
TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK
TOTAL_TOKENS = NUM_TOKENS * TOPK
TMA_LOAD_BOTH: tl.constexpr = USE_TMA_LOAD_X and USE_TMA_LOAD_dY
tidx = tl.program_id(0)
output_dtype = dW_ptr.dtype.element_ty
if USE_TMA_LOAD_dY and not TMA_LOAD_BOTH:
dY_desc = tl._experimental_make_tensor_descriptor(
dY_desc = tl.make_tensor_descriptor(
dY_ptr,
shape = [TOTAL_TOKENS, N],
strides = [N, 1],
@ -330,7 +332,7 @@ def _grouped_gemm_dW_kernel(
)
if USE_TMA_LOAD_X and not TMA_LOAD_BOTH:
x_desc = tl._experimental_make_tensor_descriptor(
x_desc = tl.make_tensor_descriptor(
x_ptr,
shape = [TOTAL_TOKENS, K],
strides = [K, 1],
@ -349,7 +351,7 @@ def _grouped_gemm_dW_kernel(
if USE_TMA_STORE:
tl.static_assert(N % BLOCK_SIZE_N == 0, "N must be divisible by BLOCK_SIZE_N")
tl.static_assert(K % BLOCK_SIZE_K == 0, "K must be divisible by BLOCK_SIZE_K")
dW_desc = tl._experimental_make_tensor_descriptor(
dW_desc = tl.make_tensor_descriptor(
dW_ptr,
shape = [NUM_EXPERTS, N, K],
strides = [N * K, K, 1],
@ -390,14 +392,14 @@ def _grouped_gemm_dW_kernel(
if m_size > 0:
if TMA_LOAD_BOTH:
dY_desc = tl._experimental_make_tensor_descriptor(
dY_desc = tl.make_tensor_descriptor(
dY_ptr,
shape = [m_end, N],
strides = [N, 1],
block_shape = [BLOCK_SIZE_M, BLOCK_SIZE_N],
)
x_desc = tl._experimental_make_tensor_descriptor(
x_desc = tl.make_tensor_descriptor(
x_ptr,
shape = [m_end, K],
strides = [K, 1],
@ -475,7 +477,7 @@ def _grouped_gemm_dW_kernel(
)
accumulator += tl.dot(
dY.T, # [BLOCK_N, BLOCK_M]
dY.T.to(x.dtype), # [BLOCK_N, BLOCK_M]
x, # [BLOCK_M, BLOCK_K]
)
@ -498,5 +500,6 @@ def _grouped_gemm_dW_kernel(
_autotuned_grouped_gemm_dW_kernel = triton.autotune(
configs = get_dW_kernel_configs(),
prune_configs_by = {"early_config_prune": prune_kernel_configs_backward_dW},
key = ["NUM_EXPERTS", "NUM_TOKENS", "N", "K", "PERMUTE_X", "PERMUTE_Y"],
# NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length
key = ["NUM_EXPERTS", "N", "K", "PERMUTE_X", "PERMUTE_Y"],
)(_grouped_gemm_dW_kernel)

View file

@ -5,7 +5,7 @@ import torch
import triton
import triton.language as tl
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
get_forward_configs,
prune_kernel_configs_fwd,
)
@ -31,11 +31,11 @@ def _grouped_gemm_forward_kernel(
topk_weights_ptr,
# Constant problem shapes
NUM_EXPERTS: tl.constexpr,
NUM_TOKENS: tl.constexpr,
NUM_TOKENS,
TOPK: tl.constexpr,
N: tl.constexpr,
K: tl.constexpr,
NUM_SMS: tl.constexpr,
NUM_SMS,
# Tuning params
BLOCK_SIZE_M: tl.constexpr,
BLOCK_SIZE_N: tl.constexpr,
@ -53,7 +53,7 @@ def _grouped_gemm_forward_kernel(
) -> None:
tl.static_assert(K % BLOCK_SIZE_K == 0)
TOTAL_TOKENS: tl.constexpr = NUM_TOKENS * TOPK
TOTAL_TOKENS = NUM_TOKENS * TOPK
SHOULD_PERMUTE: tl.constexpr = PERMUTE_X or PERMUTE_Y
SHOULD_FUSE_MUL: tl.constexpr = FUSE_MUL_PRE or FUSE_MUL_POST
SHOULD_PERMUTE_OR_FUSE: tl.constexpr = SHOULD_PERMUTE or SHOULD_FUSE_MUL
@ -66,7 +66,7 @@ def _grouped_gemm_forward_kernel(
# Also, we are defining a single global descriptor with single block shape
# Need to check that this does not result in errors when crossing expert boundaries
if USE_TMA_LOAD_X:
x_desc = tl._experimental_make_tensor_descriptor(
x_desc = tl.make_tensor_descriptor(
x_ptr,
shape = [TOTAL_TOKENS, K],
strides = [K, 1],
@ -75,7 +75,7 @@ def _grouped_gemm_forward_kernel(
if USE_TMA_LOAD_W:
expert_stride = N * K
w_desc = tl._experimental_make_tensor_descriptor(
w_desc = tl.make_tensor_descriptor(
w_ptr,
shape = [NUM_EXPERTS, N, K],
strides = [expert_stride, K, 1],
@ -100,7 +100,7 @@ def _grouped_gemm_forward_kernel(
# Need to create tma_store within loop since we need to predicate stores based on m_size
if USE_TMA_STORE:
y_desc = tl._experimental_make_tensor_descriptor(
y_desc = tl.make_tensor_descriptor(
y_ptr, # + m_start * N,
shape = [m_end, N],
strides = [N, 1],
@ -213,6 +213,7 @@ def _grouped_gemm_forward_kernel(
)
w = tl.reshape(w, (BLOCK_SIZE_N, BLOCK_SIZE_K))
x = x.to(w.dtype)
accumulator += tl.dot(x, w.T)
if not USE_TMA_LOAD_X:
@ -253,9 +254,10 @@ def _grouped_gemm_forward_kernel(
_autotuned_grouped_gemm_forward_kernel = triton.autotune(
configs = get_forward_configs(),
prune_configs_by = {"early_config_prune": prune_kernel_configs_fwd},
# NOTE: NUM_TOKENS removed from key to avoid recompilation for every sequence length
# The kernel handles variable token counts via m_sizes and tile-based processing
key = [
"NUM_EXPERTS",
"NUM_TOKENS",
"N",
"K",
"PERMUTE_X",

View file

@ -15,7 +15,7 @@ import torch
import triton
from triton.runtime.errors import OutOfResources
from grouped_gemm.kernels.autotuning import (
from .autotuning import (
BOOLS,
DEFAULT_K_BLOCK_SIZES,
DEFAULT_M_BLOCK_SIZES,

View file

@ -9,13 +9,13 @@ import torch.nn.functional as F
from transformers.models.llama4 import Llama4TextConfig
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ...interface import grouped_gemm
from ...kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from ..moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,

View file

@ -12,13 +12,13 @@ from transformers.models.qwen3_moe.modeling_qwen3_moe import (
Qwen3MoeSparseMoeBlock,
)
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ...interface import grouped_gemm
from ...kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from ..moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,

View file

@ -5,13 +5,13 @@ import torch
from transformers.models.qwen3_moe.configuration_qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeSparseMoeBlock
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
from ..interface import grouped_gemm
from ..kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from .moe_ops import (
Qwen3MoeGroupedGEMMBlock,
permute,
unpermute,

View file

@ -75,6 +75,8 @@ __all__ = [
"verify_fp8_support_if_applicable",
"_get_inference_mode_context_manager",
"hf_login",
"is_moe_model",
"get_moe_target_parameters",
"make_fast_generate_wrapper",
]
@ -599,6 +601,12 @@ try:
from transformers.configuration_utils import layer_type_validation
except:
pass
try:
# Transformers 5.0+ uses RotaryEmbeddingConfigMixin as a base class for configs
from transformers.modeling_rope_utils import RotaryEmbeddingConfigMixin
except:
pass
from transformers import __version__ as transformers_version
try:
@ -2496,6 +2504,117 @@ def hf_login(token: Optional[str] = None) -> Optional[str]:
return token
# =============================================
# MoE (Mixture of Experts) Detection and LoRA Utilities
def is_moe_model(model) -> bool:
"""
Detect if a model is a Mixture of Experts (MoE) model.
Args:
model: The model to check (can be HF model or config)
Returns:
True if the model is an MoE model, False otherwise
"""
config = getattr(model, "config", model)
# Different MoE models use different config attribute names:
# - Qwen3-MoE: num_experts
# - GLM4-MoE: n_routed_experts, num_local_experts
# - Mixtral: num_local_experts
num_experts = None
for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
num_experts = getattr(config, attr, None)
if num_experts is not None:
break
# Check text_config for VL models
if num_experts is None and hasattr(config, "text_config"):
for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
num_experts = getattr(config.text_config, attr, None)
if num_experts is not None:
break
return num_experts is not None and num_experts > 0
def get_moe_target_parameters(model, target_modules = None) -> Optional[List[str]]:
"""
Get the target_parameters for MoE expert layers if applicable.
For MoE models, returns the parameter paths for expert weights
(gate_up_proj, down_proj) that should be targeted by PEFT's
target_parameters for LoRA on nn.Parameter.
Only includes MoE parameters that match what's in target_modules:
- If "down_proj" is in target_modules -> includes "mlp.experts.down_proj"
- If "gate_proj" or "up_proj" is in target_modules -> includes "mlp.experts.gate_up_proj"
Args:
model: The model to get target parameters for
target_modules: List/tuple of target module names to match against
Returns:
List of parameter paths for MoE experts, or None if not an MoE model
"""
if not is_moe_model(model):
return None
config = getattr(model, "config", model)
# Get num_experts from various possible config attributes
num_experts = None
for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
num_experts = getattr(config, attr, None)
if num_experts is not None:
break
if num_experts is None and hasattr(config, "text_config"):
for attr in ("num_experts", "n_routed_experts", "num_local_experts"):
num_experts = getattr(config.text_config, attr, None)
if num_experts is not None:
break
if num_experts is None:
num_experts = 0
# Determine which MoE parameters to include based on target_modules
moe_params = []
# Normalize target_modules to a set for efficient lookup
if target_modules is None:
# If no target_modules specified, include all MoE params
target_set = {"gate_proj", "up_proj", "down_proj", "gate_up_proj"}
elif isinstance(target_modules, str):
target_set = {target_modules}
# Heuristic for regex matching MLPs
if "proj" in target_modules and (
"mlp" in target_modules or "ffn" in target_modules
):
target_set.update({"gate_proj", "up_proj", "down_proj", "gate_up_proj"})
else:
target_set = set(target_modules) if target_modules else set()
# gate_up_proj combines both gate_proj and up_proj in MoE
# Also match "gate_up_proj" directly since users may specify the fused name
if (
"gate_proj" in target_set
or "up_proj" in target_set
or "gate_up_proj" in target_set
):
moe_params.append("mlp.experts.gate_up_proj")
if "down_proj" in target_set:
moe_params.append("mlp.experts.down_proj")
if moe_params:
print(
f"Unsloth: Detected MoE model with {num_experts} experts - enabling LoRA on: {moe_params}"
)
return moe_params
return None
def make_fast_generate_wrapper(original_generate):
"""
Creates a wrapper around model.generate that checks for incorrect

450
unsloth/models/glm4_moe.py Normal file
View file

@ -0,0 +1,450 @@
# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
GLM-4.7 Flash (GLM4 MoE Lite) optimized implementation using grouped GEMM.
Key architecture differences from Qwen3 MoE:
- Router uses sigmoid activation (not softmax)
- Has routed_scaling_factor of 1.8
- Has 1 shared expert that processes all tokens
- Uses group-based selection before topk
- Uses MLA (Multi-head Latent Attention)
"""
from .llama import *
import os
from ._utils import __version__
from .llama import (
LlamaRotaryEmbedding,
LlamaLinearScalingRotaryEmbedding,
fix_prepare_inputs_for_generation,
fast_rms_layernorm_inference,
fast_swiglu_inference,
LlamaModel_fast_forward,
LlamaModel_fast_forward_inference,
CausalLM_fast_forward,
PeftModel_fast_forward,
)
import torch
import torch.nn.functional as F
from typing import Optional, Tuple
from ..kernels import fast_rms_layernorm
# Import the grouped gemm utilities from unsloth kernels
# The grouped_gemm module expects its parent directory to be in sys.path
HAS_GROUPED_GEMM = False
try:
import sys
import os
# Add the moe directory (parent of grouped_gemm) to sys.path
_moe_path = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "kernels", "moe"
)
if _moe_path not in sys.path:
sys.path.insert(0, _moe_path)
# Import grouped_gemm package first to apply TMA compatibility shim
# This patches triton.language to support both old and new TMA API names
import grouped_gemm # noqa: F401 - triggers TMA compatibility shim
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.reference.moe_ops import (
get_routing_indices,
permute,
unpermute,
)
HAS_GROUPED_GEMM = True
except ImportError as e:
import warnings
warnings.warn(
f"Grouped GEMM not available: {e}. MoE will use fallback implementation."
)
# Import transformers GLM4 MoE Lite classes
try:
from transformers.models.glm4_moe_lite.modeling_glm4_moe_lite import (
Glm4MoeLiteAttention,
Glm4MoeLiteMoE,
Glm4MoeLiteMLP,
Glm4MoeLiteNaiveMoe,
Glm4MoeLiteTopkRouter,
Glm4MoeLiteDecoderLayer,
Glm4MoeLiteModel,
Glm4MoeLiteForCausalLM,
Glm4MoeLiteRMSNorm,
)
HAS_GLM4_MOE = True
except ImportError:
HAS_GLM4_MOE = False
# Create dummy classes for type checking
class Glm4MoeLiteAttention:
pass
class Glm4MoeLiteMoE:
pass
class Glm4MoeLiteMLP:
pass
class Glm4MoeLiteNaiveMoe:
pass
class Glm4MoeLiteTopkRouter:
pass
class Glm4MoeLiteDecoderLayer:
pass
class Glm4MoeLiteModel:
pass
class Glm4MoeLiteForCausalLM:
pass
torch_nn_functional_silu = torch.nn.functional.silu
def Glm4MoeLiteMoE_fast_forward(self, hidden_states):
"""
Optimized MoE forward pass using grouped GEMM.
GLM4 MoE specifics:
- Uses sigmoid router activation (not softmax)
- Has routed_scaling_factor of 1.8
- Has 1 shared expert that always processes all tokens
- Uses group-based selection with topk_group
"""
residuals = hidden_states
orig_shape = hidden_states.shape
batch_size, seq_len, hidden_dim = orig_shape
num_tokens = batch_size * seq_len
# Flatten hidden states for routing
hidden_states = hidden_states.view(-1, hidden_dim)
# Router computation
router_logits = self.gate(hidden_states) # [num_tokens, n_routed_experts]
topk_indices, topk_weights = self.route_tokens_to_experts(router_logits)
# Cast routing weights to match hidden_states dtype (Qwen3 pattern)
# Sigmoid router returns fp32, but hidden_states may be bf16
topk_weights = topk_weights.to(hidden_states.dtype)
# Get routing indices for grouped GEMM
with torch.no_grad():
token_counts_by_expert, gather_indices = get_routing_indices(
topk_indices, self.n_routed_experts
)
# Use grouped GEMM for expert computation
if HAS_GROUPED_GEMM:
# Cast hidden_states to match expert weights dtype
# Under autocast, hidden_states may be fp32 while weights are bf16
hidden_states = hidden_states.to(self.experts.gate_up_proj.dtype)
# First grouped GEMM: gate_up_proj with permute_x
# Input: [num_tokens, hidden_dim] -> Output: [total_tokens, 2*intermediate_dim]
intermediate = grouped_gemm(
X = hidden_states,
W = self.experts.gate_up_proj,
m_sizes = token_counts_by_expert.int(),
topk = self.top_k,
gather_indices = gather_indices,
permute_x = True,
permute_y = False,
autotune = True,
is_first_gemm = True,
)
# Activation: SiLU(gate) * up
gate, up = intermediate.chunk(2, dim = -1)
intermediate = torch_nn_functional_silu(gate) * up
# Second grouped GEMM: down_proj with permute_y
# Input: [total_tokens, intermediate_dim] -> Output: [total_tokens, hidden_dim]
expert_output = grouped_gemm(
X = intermediate,
W = self.experts.down_proj,
m_sizes = token_counts_by_expert.int(),
topk = self.top_k,
gather_indices = gather_indices,
permute_x = False,
permute_y = True,
autotune = True,
is_first_gemm = False,
)
# Merge topk weights: [num_tokens, top_k, hidden_dim] -> [num_tokens, hidden_dim]
hidden_states = (
expert_output.view(num_tokens, self.top_k, hidden_dim)
* topk_weights.unsqueeze(-1)
).sum(dim = 1)
else:
# Fallback to naive implementation
hidden_states = self.experts(hidden_states, topk_indices, topk_weights)
# Add shared expert output
hidden_states = hidden_states + self.shared_experts(residuals.view(-1, hidden_dim))
return hidden_states.view(*orig_shape)
def Glm4MoeLiteNaiveMoe_fast_forward(
self,
hidden_states: torch.Tensor,
top_k_index: torch.Tensor,
top_k_weights: torch.Tensor,
) -> torch.Tensor:
"""
Optimized expert forward using grouped GEMM.
Args:
hidden_states: [num_tokens, hidden_dim]
top_k_index: [num_tokens, top_k] indices of selected experts
top_k_weights: [num_tokens, top_k] weights for selected experts
Returns:
[num_tokens, hidden_dim] output after weighted sum of expert outputs
"""
num_tokens, hidden_dim = hidden_states.shape
top_k = top_k_index.shape[1]
# Cast routing weights to match hidden_states dtype (Qwen3 pattern)
top_k_weights = top_k_weights.to(hidden_states.dtype)
if not HAS_GROUPED_GEMM:
# Fallback to original naive implementation
final_hidden_states = torch.zeros_like(hidden_states)
with torch.no_grad():
expert_mask = torch.nn.functional.one_hot(
top_k_index, num_classes = self.num_experts
)
expert_mask = expert_mask.permute(2, 1, 0)
expert_hit = torch.greater(expert_mask.sum(dim = (-1, -2)), 0).nonzero()
for expert_idx in expert_hit:
expert_idx = expert_idx[0]
if expert_idx == self.num_experts:
continue
top_k_pos, token_idx = torch.where(expert_mask[expert_idx])
current_state = hidden_states[token_idx]
gate, up = torch.nn.functional.linear(
current_state, self.gate_up_proj[expert_idx]
).chunk(2, dim = -1)
current_hidden_states = self.act_fn(gate) * up
current_hidden_states = torch.nn.functional.linear(
current_hidden_states, self.down_proj[expert_idx]
)
current_hidden_states = (
current_hidden_states * top_k_weights[token_idx, top_k_pos, None]
)
final_hidden_states.index_add_(
0, token_idx, current_hidden_states.to(final_hidden_states.dtype)
)
return final_hidden_states
# Get routing indices for grouped GEMM
with torch.no_grad():
token_counts_by_expert, gather_indices = get_routing_indices(
top_k_index, self.num_experts
)
# Cast hidden_states to match expert weights dtype
# Under autocast, hidden_states may be fp32 while weights are bf16
hidden_states = hidden_states.to(self.gate_up_proj.dtype)
# First grouped GEMM: gate_up_proj
intermediate = grouped_gemm(
X = hidden_states,
W = self.gate_up_proj,
m_sizes = token_counts_by_expert.int(),
topk = top_k,
gather_indices = gather_indices,
permute_x = True,
permute_y = False,
autotune = True,
is_first_gemm = True,
)
# Activation: SiLU(gate) * up
gate, up = intermediate.chunk(2, dim = -1)
intermediate = self.act_fn(gate) * up
# Second grouped GEMM: down_proj
expert_output = grouped_gemm(
X = intermediate,
W = self.down_proj,
m_sizes = token_counts_by_expert.int(),
topk = top_k,
gather_indices = gather_indices,
permute_x = False,
permute_y = True,
autotune = True,
is_first_gemm = False,
)
# Merge topk weights
final_hidden_states = (
expert_output.view(num_tokens, top_k, hidden_dim) * top_k_weights.unsqueeze(-1)
).sum(dim = 1)
return final_hidden_states
def Glm4MoeLiteDecoderLayer_fast_forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
past_key_values = None,
use_cache: bool = False,
cache_position: Optional[torch.LongTensor] = None,
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
**kwargs,
) -> torch.Tensor:
"""
Optimized decoder layer forward with fast RMS layernorm.
"""
# Check if we're in inference mode
is_inference = use_cache and hasattr(self, "_flag_for_generation")
if is_inference:
# Self-attention with fast inference path
residual = hidden_states
hidden_states = fast_rms_layernorm_inference(
self.input_layernorm, hidden_states
)
hidden_states, _ = self.self_attn(
hidden_states = hidden_states,
attention_mask = attention_mask,
position_ids = position_ids,
past_key_values = past_key_values,
use_cache = use_cache,
cache_position = cache_position,
position_embeddings = position_embeddings,
**kwargs,
)
hidden_states = residual + hidden_states
# MLP/MoE
residual = hidden_states
hidden_states = fast_rms_layernorm_inference(
self.post_attention_layernorm, hidden_states
)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
else:
# Training path
residual = hidden_states
hidden_states = fast_rms_layernorm(self.input_layernorm, hidden_states)
hidden_states, _ = self.self_attn(
hidden_states = hidden_states,
attention_mask = attention_mask,
position_ids = position_ids,
past_key_values = past_key_values,
use_cache = use_cache,
cache_position = cache_position,
position_embeddings = position_embeddings,
**kwargs,
)
hidden_states = residual + hidden_states
# MLP/MoE
residual = hidden_states
hidden_states = fast_rms_layernorm(self.post_attention_layernorm, hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
return hidden_states
def Glm4MoeLiteMLP_fast_forward(self, x):
"""
Optimized MLP forward using fused SwiGLU.
"""
return fast_swiglu_inference(self, x)
class FastGLM47Model(FastLlamaModel):
"""
Fast GLM-4.7 Flash (GLM4 MoE Lite) model with grouped GEMM optimization.
This provides 2-3x throughput improvement for MoE layers by:
- Replacing sequential expert loops with grouped GEMM operations
- Fusing permutation operations into the GEMM kernels
- Using optimized RMS LayerNorm and SwiGLU implementations
"""
@staticmethod
def pre_patch():
if not HAS_GLM4_MOE:
raise ImportError(
"Unsloth: GLM4 MoE Lite support requires transformers >= 5.0.0. "
"Please upgrade with: pip install --upgrade transformers"
)
# Patch MoE forward with grouped GEMM optimization
# TMA compatibility is handled by grouped_gemm/__init__.py which patches
# triton.language to support both old (_experimental_make_tensor_descriptor)
# and new (make_tensor_descriptor) API names
if HAS_GROUPED_GEMM:
Glm4MoeLiteNaiveMoe.forward = Glm4MoeLiteNaiveMoe_fast_forward
Glm4MoeLiteMoE.forward = Glm4MoeLiteMoE_fast_forward
# Note: We don't patch the following for GLM4 MoE because:
# - GLM4 uses MLA (Multi-head Latent Attention) which has different projection names
# - Glm4MoeLiteRotaryEmbedding doesn't have extend_rope_embedding method
# - The decoder layer and model forward functions assume Llama-compatible infrastructure
return
@staticmethod
def from_pretrained(
model_name = "unsloth/GLM-4.7-Flash",
max_seq_length = 4096,
dtype = None,
load_in_4bit = True,
token = None,
device_map = "sequential",
rope_scaling = None,
fix_tokenizer = True,
model_patcher = None,
tokenizer_name = None,
trust_remote_code = False,
**kwargs,
):
# Pop kwargs that are used by loader but not passed to model
kwargs.pop("unsloth_force_compile", None)
return FastLlamaModel.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
dtype = dtype,
load_in_4bit = load_in_4bit,
token = token,
device_map = device_map,
rope_scaling = rope_scaling,
fix_tokenizer = fix_tokenizer,
model_patcher = FastGLM47Model,
tokenizer_name = tokenizer_name,
trust_remote_code = trust_remote_code,
**kwargs,
)

View file

@ -2659,6 +2659,7 @@ class FastLlamaModel:
loftq_config = {},
temporary_location = "_unsloth_temporary_saved_buffers",
qat_scheme = None,
target_parameters = None, # For MoE expert layers (nn.Parameter)
ensure_weight_tying = False,
**kwargs,
):
@ -2689,6 +2690,7 @@ class FastLlamaModel:
init_lora_weights = init_lora_weights,
loftq_config = loftq_config,
temporary_location = temporary_location,
target_parameters = target_parameters,
ensure_weight_tying = ensure_weight_tying,
**kwargs,
)
@ -2974,6 +2976,10 @@ class FastLlamaModel:
# Does not get lora yet, so get name from model, not base model
is_classification = "Classification" in str(type(model))
# Auto-detect MoE models and populate target_parameters for expert layers
if target_parameters is None:
target_parameters = get_moe_target_parameters(model, target_modules)
arguments = dict(
r = r,
lora_alpha = lora_alpha,
@ -2986,6 +2992,7 @@ class FastLlamaModel:
loftq_config = loftq_config,
use_rslora = use_rslora,
modules_to_save = modules_to_save,
target_parameters = target_parameters,
ensure_weight_tying = ensure_weight_tying,
**kwargs,
)

View file

@ -736,6 +736,7 @@ class FastModel(FastBaseModel):
qat_scheme = None,
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
unsloth_tiled_mlp = False,
target_parameters = None, # For MoE expert parameters
*args,
**kwargs,
):

View file

@ -98,6 +98,7 @@ VLLM_SUPPORTED_VLM = [
"gemma3",
"mistral3",
"qwen3_vl",
"qwen3_vl_moe",
]
VLLM_NON_LORA_VLM = [
"mllama",
@ -960,6 +961,7 @@ class FastBaseModel:
task_type = TaskType.CAUSAL_LM,
temporary_location = "_unsloth_temporary_saved_buffers",
qat_scheme = None,
target_parameters = None, # For MoE expert layers (nn.Parameter)
ensure_weight_tying = False, # [TODO] Add `ensure_weight_tying` for `modules_to_save` for vision models
**kwargs,
):
@ -1041,6 +1043,10 @@ class FastBaseModel:
loftq_config, lora_dropout, bias, init_lora_weights, model
)
# Auto-detect MoE models and populate target_parameters for expert layers
if target_parameters is None:
target_parameters = get_moe_target_parameters(model, target_modules)
# Get only allowed parameters for LoraConfig
local_variables = {
**locals(),