Merge branch 'main' of https://github.com/unslothai/unsloth
This commit is contained in:
commit
5000413815
14 changed files with 1298 additions and 95 deletions
500
unsloth/kernels/moe/autotune_cache.py
Normal file
500
unsloth/kernels/moe/autotune_cache.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
450
unsloth/models/glm4_moe.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue