Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,8 @@ checkstyle:
# We have to explicitly set HF_DATASETS_OFFLINE=1, or dataset will silently try to send metrics and timeout (80s) https://github.com/huggingface/datasets/blob/37a603679f451826cfafd8aae00738b01dcb9d58/src/datasets/load.py#L286
test-convergence:
HF_DATASETS_OFFLINE=1 python -m pytest --disable-warnings \
test/convergence/fp32/test_lfm2_models.py \
test/convergence/bf16/test_lfm2_models.py \
test/convergence/fp32/test_mini_models.py \
test/convergence/fp32/test_mini_models_multimodal.py \
test/convergence/fp32/test_mini_models_with_logits.py \
Expand Down
5 changes: 5 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -297,6 +297,9 @@ loss.backward()
| Llama4 (Text) & (Multimodal) | `liger_kernel.transformers.apply_liger_kernel_to_llama4` | RMSNorm, LayerNorm, GeGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| LLaMA 2 & 3 | `liger_kernel.transformers.apply_liger_kernel_to_llama` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| LLaMA 3.2-Vision | `liger_kernel.transformers.apply_liger_kernel_to_mllama` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2 | `liger_kernel.transformers.apply_liger_kernel_to_lfm2` | RoPE, RMSNorm, SwiGLU, ShortConv, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2MoE | `liger_kernel.transformers.apply_liger_kernel_to_lfm2_moe` | RoPE, RMSNorm, SwiGLU, ShortConv, FusedMoE, MoERouter, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2VL | `liger_kernel.transformers.apply_liger_kernel_to_lfm2_vl` | SigLIP2 LayerNorm, RoPE, RMSNorm, SwiGLU, ShortConv, CrossEntropyLoss, FusedLinearCrossEntropy |
| Ministral | `liger_kernel.transformers.apply_liger_kernel_to_ministral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| Mistral | `liger_kernel.transformers.apply_liger_kernel_to_mistral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| Mixtral | `liger_kernel.transformers.apply_liger_kernel_to_mixtral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
Expand Down Expand Up @@ -346,6 +349,8 @@ loss.backward()
| CrossEntropy | `liger_kernel.transformers.LigerCrossEntropyLoss` |
| Fused Linear CrossEntropy | `liger_kernel.transformers.LigerFusedLinearCrossEntropyLoss`|
| Multi Token Attention | `liger_kernel.transformers.LigerMultiTokenAttention` |
| LFM2 Short Convolution | `liger_kernel.ops.LigerLfm2ShortConvFunction` |
| LFM2 MoE Router | `liger_kernel.ops.LigerLfm2MoeRouterFunction` |
| Softmax | `liger_kernel.transformers.LigerSoftmax` |
| Sparsemax | `liger_kernel.transformers.LigerSparsemax` |
| mHC (Hyper-Connections) | `liger_kernel.transformers.LigerMHC` |
Expand Down
336 changes: 336 additions & 0 deletions benchmark/data/all_benchmark_data.csv

Large diffs are not rendered by default.

116 changes: 116 additions & 0 deletions benchmark/scripts/benchmark_lfm2_moe_router.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
import torch
import torch.nn as nn

from benchmark_model_configs import MODEL_REGISTRY
from benchmark_model_configs import build_model_config_sweep
from benchmark_model_configs import build_token_length_sweep
from benchmark_model_configs import get_benchmark_model_config
from test.transformers.test_lfm2_moe_router import _reference
from utils import SingleBenchmarkRunInput
from utils import build_memory_bench_fn
from utils import build_speed_bench_fn
from utils import parse_benchmark_script_args
from utils import run_benchmarks

from liger_kernel.ops import LigerLfm2MoeRouterFunction
from liger_kernel.utils import infer_device

device = infer_device()


class _Router(nn.Module):
def __init__(self, num_experts, top_k, dtype, use_liger):
super().__init__()
self.register_buffer("expert_bias", torch.zeros(num_experts, device=device, dtype=torch.float32))
self.top_k = top_k
self.use_liger = use_liger
self.dtype = dtype

def forward(self, router_logits):
if self.use_liger:
return LigerLfm2MoeRouterFunction.apply(
router_logits,
self.expert_bias,
self.top_k,
True,
1.0,
)[1]
return _reference(router_logits, self.expert_bias, self.top_k, True, 1.0)[1]


def setup_lfm2_moe_router(input: SingleBenchmarkRunInput):
cfg = input.extra_benchmark_config
if isinstance(input.x, str):
model_cfg = MODEL_REGISTRY[input.x]
num_tokens = cfg["bsz"] * cfg["seq_len"]
num_experts = model_cfg.num_experts
top_k = model_cfg.topk
dtype = model_cfg.dtype
else:
num_tokens = cfg["bsz"] * input.x
num_experts = cfg["num_experts"]
top_k = cfg["topk"]
dtype = cfg["dtype"]

if num_experts is None or top_k is None:
raise ValueError("LFM2 MoE router benchmarks require an MoE model configuration")

router_logits = torch.randn(
num_tokens,
num_experts,
device=device,
dtype=dtype,
requires_grad=True,
)
if input.kernel_provider == "liger":
layer = _Router(num_experts, top_k, dtype, use_liger=True)
elif input.kernel_provider == "huggingface":
layer = _Router(num_experts, top_k, dtype, use_liger=False)
else:
raise ValueError(f"Invalid provider: {input.kernel_provider} for LFM2 MoE router")
return router_logits, layer


if __name__ == "__main__":
args = parse_benchmark_script_args()

if args.sweep_mode == "model_config":
moe_configs = [config for config in MODEL_REGISTRY.values() if config.is_moe]
common_configs = build_model_config_sweep(
kernel_name="lfm2_moe_router",
all_model_configs=moe_configs,
setup_fn=setup_lfm2_moe_router,
model_keys=["num_experts", "topk", "dtype"],
probe_provider="huggingface",
extra_configs={"bsz": 1},
probe_dim="T",
bt=args.bt,
overwrite=args.overwrite,
)
else:
model = get_benchmark_model_config(args.model or "lfm2_moe_8b_a1b")
common_configs = build_token_length_sweep(
kernel_name="lfm2_moe_router",
probe_x=1024,
model=model,
setup_fn=setup_lfm2_moe_router,
model_keys=["num_experts", "topk", "dtype"],
extra_configs={"bsz": 1},
scale_dim="T",
x_label="total tokens",
probe_provider="huggingface",
overwrite=args.overwrite,
)

common_configs["kernel_providers"] = ["huggingface", "liger"]
for metric_name, metric_unit, bench_fn in (
("speed", "ms", build_speed_bench_fn(setup_lfm2_moe_router)),
("memory", "MB", build_memory_bench_fn(setup_lfm2_moe_router)),
):
run_benchmarks(
bench_test_fn=bench_fn,
kernel_operation_modes=["forward", "backward", "full"],
metric_name=metric_name,
metric_unit=metric_unit,
**common_configs,
)
103 changes: 103 additions & 0 deletions benchmark/scripts/benchmark_lfm2_short_conv.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
import torch
import torch.nn as nn

from benchmark_model_configs import MODEL_REGISTRY
from benchmark_model_configs import build_model_config_sweep
from benchmark_model_configs import build_token_length_sweep
from benchmark_model_configs import get_benchmark_model_config
from test.transformers.test_lfm2_short_conv import _reference
from utils import SingleBenchmarkRunInput
from utils import build_memory_bench_fn
from utils import build_speed_bench_fn
from utils import parse_benchmark_script_args
from utils import run_benchmarks

from liger_kernel.ops import LigerLfm2ShortConvFunction
from liger_kernel.utils import infer_device

device = infer_device()


class _ShortConv(nn.Module):
def __init__(self, hidden_size, kernel_size, dtype, use_liger, bias_enabled):
super().__init__()
self.weight = nn.Parameter(torch.randn(hidden_size, 1, kernel_size, device=device, dtype=dtype) * 0.02)
self.bias = nn.Parameter(torch.zeros(hidden_size, device=device, dtype=dtype)) if bias_enabled else None
self.use_liger = use_liger

def forward(self, bcx):
if self.use_liger:
return LigerLfm2ShortConvFunction.apply(bcx, self.weight, self.bias)
return _reference(bcx, self.weight, self.bias)


def setup_lfm2_short_conv(input: SingleBenchmarkRunInput):
cfg = input.extra_benchmark_config
if isinstance(input.x, str):
model_cfg = MODEL_REGISTRY[input.x]
seq_len = cfg["seq_len"]
hidden_size = model_cfg.hidden_size
dtype = model_cfg.dtype
else:
seq_len = input.x
hidden_size = cfg["hidden_size"]
dtype = cfg["dtype"]

bcx = torch.randn(
cfg["bsz"],
seq_len,
3 * hidden_size,
device=device,
dtype=dtype,
requires_grad=True,
)
if input.kernel_provider == "liger":
layer = _ShortConv(hidden_size, cfg["kernel_size"], dtype, use_liger=True, bias_enabled=cfg["bias_enabled"])
elif input.kernel_provider == "huggingface":
layer = _ShortConv(hidden_size, cfg["kernel_size"], dtype, use_liger=False, bias_enabled=cfg["bias_enabled"])
else:
raise ValueError(f"Invalid provider: {input.kernel_provider} for LFM2 short convolution")
return bcx, layer


if __name__ == "__main__":
args = parse_benchmark_script_args()

if args.sweep_mode == "model_config":
common_configs = build_model_config_sweep(
kernel_name="lfm2_short_conv",
setup_fn=setup_lfm2_short_conv,
model_keys=["hidden_size", "dtype"],
probe_provider="huggingface",
extra_configs={"bsz": 1, "kernel_size": 3, "bias_enabled": False},
probe_dim="T",
bt=args.bt,
overwrite=args.overwrite,
)
else:
model = get_benchmark_model_config(args.model or "lfm2_1.2b")
common_configs = build_token_length_sweep(
kernel_name="lfm2_short_conv",
probe_x=1024,
model=model,
setup_fn=setup_lfm2_short_conv,
model_keys=["hidden_size", "dtype"],
extra_configs={"bsz": 1, "kernel_size": 3, "bias_enabled": False},
scale_dim="T",
x_label="total tokens",
probe_provider="huggingface",
overwrite=args.overwrite,
)

common_configs["kernel_providers"] = ["huggingface", "liger"]
for metric_name, metric_unit, bench_fn in (
("speed", "ms", build_speed_bench_fn(setup_lfm2_short_conv)),
("memory", "MB", build_memory_bench_fn(setup_lfm2_short_conv)),
):
run_benchmarks(
bench_test_fn=bench_fn,
kernel_operation_modes=["forward", "backward", "full"],
metric_name=metric_name,
metric_unit=metric_unit,
**common_configs,
)
29 changes: 29 additions & 0 deletions benchmark/scripts/benchmark_model_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,33 @@ class MoEModelConfig:
topk=8,
)

LFM2_1_2B = ModelConfig(
name="lfm2_1.2b",
hidden_size=2048,
intermediate_size=12288,
vocab_size=65536,
num_attention_heads=32,
num_key_value_heads=8,
head_dim=64,
hidden_act="silu",
max_position_embeddings=128000,
)

LFM2_MOE_8B_A1B = ModelConfig(
name="lfm2_moe_8b_a1b",
hidden_size=2048,
intermediate_size=7168,
vocab_size=65536,
num_attention_heads=32,
num_key_value_heads=8,
head_dim=64,
hidden_act="silu",
max_position_embeddings=128000,
num_experts=32,
topk=4,
moe_intermediate_size=1792,
)

MODEL_REGISTRY: Dict[str, ModelConfig] = {
"llama_2_7b": LLAMA_2_7B,
"llama_3_8b": LLAMA_3_8B,
Expand All @@ -231,6 +258,8 @@ class MoEModelConfig:
"qwen2.5_72b": QWEN_2_5_72B,
"deepseek_v2_lite": DEEPSEEK_V2_LITE,
"deepseek_v3": DEEPSEEK_V3,
"lfm2_1.2b": LFM2_1_2B,
"lfm2_moe_8b_a1b": LFM2_MOE_8B_A1B,
}

DEFAULT_MODEL_CONFIG = LLAMA_3_8B
Expand Down
21 changes: 21 additions & 0 deletions docs/High-Level-APIs.md
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,9 @@ You can also use the Patching APIs to use the kernels for a specific model archi
|-------------|--------------------------------------------------------------|-------------------------------------------------------------------------|
| LLaMA 2 & 3 | `liger_kernel.transformers.apply_liger_kernel_to_llama` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| LLaMA 3.2-Vision | `liger_kernel.transformers.apply_liger_kernel_to_mllama` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2 | `liger_kernel.transformers.apply_liger_kernel_to_lfm2` | RoPE, RMSNorm, SwiGLU, ShortConv, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2MoE | `liger_kernel.transformers.apply_liger_kernel_to_lfm2_moe` | RoPE, RMSNorm, SwiGLU, ShortConv, FusedMoE, MoERouter, CrossEntropyLoss, FusedLinearCrossEntropy |
| LFM2VL | `liger_kernel.transformers.apply_liger_kernel_to_lfm2_vl` | SigLIP2 LayerNorm, RoPE, RMSNorm, SwiGLU, ShortConv, CrossEntropyLoss, FusedLinearCrossEntropy |
| Mistral | `liger_kernel.transformers.apply_liger_kernel_to_mistral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| Mixtral | `liger_kernel.transformers.apply_liger_kernel_to_mixtral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| Gemma1 | `liger_kernel.transformers.apply_liger_kernel_to_gemma` | RoPE, RMSNorm, GeGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
Expand Down Expand Up @@ -62,6 +65,24 @@ You can also use the Patching APIs to use the kernels for a specific model archi
show_docstring: true
show_signature: true

::: liger_kernel.transformers.apply_liger_kernel_to_lfm2
options:
extra:
show_docstring: true
show_signature: true

::: liger_kernel.transformers.apply_liger_kernel_to_lfm2_moe
options:
extra:
show_docstring: true
show_signature: true

::: liger_kernel.transformers.apply_liger_kernel_to_lfm2_vl
options:
extra:
show_docstring: true
show_signature: true

::: liger_kernel.transformers.apply_liger_kernel_to_gemma
options:
extra:
Expand Down
2 changes: 2 additions & 0 deletions src/liger_kernel/ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,8 @@
from liger_kernel.ops.layer_norm import LigerLayerNormFunction # noqa: F401
from liger_kernel.ops.layer_norm import layer_norm_backward # noqa: F401
from liger_kernel.ops.layer_norm import layer_norm_forward # noqa: F401
from liger_kernel.ops.lfm2_moe_router import LigerLfm2MoeRouterFunction # noqa: F401
from liger_kernel.ops.lfm2_short_conv import LigerLfm2ShortConvFunction # noqa: F401
from liger_kernel.ops.llama4_rope import LigerLlama4RopeFunction # noqa: F401
from liger_kernel.ops.mhc import LigerMHCCoeffsFunction # noqa: F401
from liger_kernel.ops.mhc import LigerMHCPostResFunction # noqa: F401
Expand Down
Loading