Skip to content
Open
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -387,6 +387,7 @@ loss.backward()
| OLMo2 | `liger_kernel.transformers.apply_liger_kernel_to_olmo2` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| Olmo3 | `liger_kernel.transformers.apply_liger_kernel_to_olmo3` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| GLM-4 | `liger_kernel.transformers.apply_liger_kernel_to_glm4` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
| DeepSeek-V3 | `liger_kernel.transformers.apply_liger_kernel_to_deepseek_v3` | RMSNorm, SwiGLU (dense/shared MLPs), CrossEntropyLoss, FusedLinearCrossEntropy |
| DeepSeek-V4 | `liger_kernel.transformers.apply_liger_kernel_to_deepseek_v4` | RMSNorm, CrossEntropyLoss, FusedLinearCrossEntropy |
| GPT-OSS | `liger_kernel.transformers.apply_liger_kernel_to_gpt_oss` | RoPE, RMSNorm, CrossEntropyLoss, FusedLinearCrossEntropy |
| InternVL3 | `liger_kernel.transformers.apply_liger_kernel_to_internvl` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy |
Expand Down
3 changes: 3 additions & 0 deletions src/liger_kernel/transformers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
from liger_kernel.transformers.auto_model import AutoLigerKernelForCausalLM # noqa: F401
from liger_kernel.transformers.monkey_patch import _apply_liger_kernel # noqa: F401
from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance # noqa: F401
from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_deepseek_v3 # noqa: F401
from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_deepseek_v4 # noqa: F401
from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_exaone4 # noqa: F401
from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_falcon_h1 # noqa: F401
Expand Down Expand Up @@ -158,6 +159,7 @@ def __getattr__(name: str):
"apply_liger_kernel_to_smolvlm",
"apply_liger_kernel_to_hunyuan_v1_dense",
"apply_liger_kernel_to_hunyuan_v1_moe",
"apply_liger_kernel_to_deepseek_v3",
"apply_liger_kernel_to_deepseek_v4",
"apply_liger_kernel_to_exaone4",
}
Expand Down Expand Up @@ -249,6 +251,7 @@ def __getattr__(name: str):
"apply_liger_kernel_to_smolvlm",
"apply_liger_kernel_to_hunyuan_v1_dense",
"apply_liger_kernel_to_hunyuan_v1_moe",
"apply_liger_kernel_to_deepseek_v3",
"apply_liger_kernel_to_deepseek_v4",
"apply_liger_kernel_to_exaone4",
]
Expand Down
81 changes: 81 additions & 0 deletions src/liger_kernel/transformers/model/deepseek_v3.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
import torch

from transformers.cache_utils import Cache
from transformers.modeling_outputs import BaseModelOutputWithPast
from transformers.utils import can_return_tuple

from liger_kernel.transformers.model.llama import lce_maybe_trainable_lm_head
from liger_kernel.transformers.model.loss_utils import unpack_cross_entropy_result
from liger_kernel.transformers.model.output_classes import LigerCausalLMOutputWithPast


@can_return_tuple
def lce_forward(
self,
input_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
past_key_values: Cache | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = None,
logits_to_keep: int | torch.Tensor = 0,
skip_logits: bool | None = None,
**kwargs,
) -> LigerCausalLMOutputWithPast:
outputs: BaseModelOutputWithPast = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
past_key_values=past_key_values,
inputs_embeds=inputs_embeds,
use_cache=use_cache,
**kwargs,
)

hidden_states = outputs.last_hidden_state
slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
kept_hidden_states = hidden_states[:, slice_indices, :]

shift_labels = kwargs.pop("shift_labels", None)
logits = None
loss = None
token_accuracy = None
predicted_tokens = None

if skip_logits and labels is None and shift_labels is None:
raise ValueError("skip_logits is True, but labels and shift_labels are None")

if skip_logits is None:
skip_logits = self.training and (labels is not None or shift_labels is not None)

if skip_logits:
result = lce_maybe_trainable_lm_head(
self,
hidden_states=kept_hidden_states,
hidden_size=self.config.hidden_size,
labels=labels,
shift_labels=shift_labels,
**kwargs,
)
loss, _, token_accuracy, predicted_tokens = unpack_cross_entropy_result(result)
else:
logits = self.lm_head(kept_hidden_states)
if labels is not None or shift_labels is not None:
loss = self.loss_function(
logits=logits,
labels=labels,
shift_labels=shift_labels,
vocab_size=self.config.vocab_size,
**kwargs,
)

return LigerCausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=outputs.past_key_values,
hidden_states=outputs.hidden_states,
attentions=outputs.attentions,
token_accuracy=token_accuracy,
predicted_tokens=predicted_tokens,
)
85 changes: 85 additions & 0 deletions src/liger_kernel/transformers/monkey_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -3377,6 +3377,90 @@ def apply_liger_kernel_to_hunyuan_v1_moe(
_patch_rms_norm_module(decoder_layer.post_attention_layernorm)


def apply_liger_kernel_to_deepseek_v3(
rope: bool = False,
cross_entropy: bool = False,
fused_linear_cross_entropy: bool = True,
rms_norm: bool = True,
swiglu: bool = True,
model: PreTrainedModel = None,
) -> None:
"""
Apply Liger kernels to replace original implementation in HuggingFace DeepSeek-V3 models.

NOTE: RoPE is not supported for DeepSeek-V3. Its attention uses interleaved partial RoPE,
which is incompatible with ``liger_rotary_pos_emb``. Routed experts are intentionally left
unchanged; SwiGLU is only applied to dense and shared-expert MLPs.

Args:
rope (bool): Whether to apply Liger's rotary position embedding. Default is False.
Currently unsupported; emits a warning and is a no-op.
cross_entropy (bool): Whether to apply Liger's cross entropy loss. Default is False.
fused_linear_cross_entropy (bool):
Whether to apply Liger's fused linear cross entropy loss. Default is True.
`cross_entropy` and `fused_linear_cross_entropy` cannot both be True.
If `fused_linear_cross_entropy` is True, the logits will not be materialized but more memory efficient.
rms_norm (bool): Whether to apply Liger's RMSNorm. Default is True.
swiglu (bool): Whether to apply Liger's SwiGLU to dense and shared-expert MLPs. Default is True.
model (PreTrainedModel): The model instance to apply Liger kernels to, if already loaded.
Default is None.
"""
assert not (cross_entropy and fused_linear_cross_entropy), (
"cross_entropy and fused_linear_cross_entropy cannot both be True."
)

from transformers.models.deepseek_v3 import modeling_deepseek_v3
from transformers.models.deepseek_v3.modeling_deepseek_v3 import DeepseekV3Model

from liger_kernel.transformers.model.deepseek_v3 import lce_forward as deepseek_v3_lce_forward
from liger_kernel.transformers.swiglu import LigerQwen3MoeSwiGLUMLP

if rope:
logger.warning_once(
"rope=True is not supported for DeepSeek-V3: interleaved partial RoPE is "
"incompatible with liger_rotary_pos_emb. Skipping rope kernel swap."
)

if rms_norm:
modeling_deepseek_v3.DeepseekV3RMSNorm = LigerRMSNorm

if cross_entropy:
from transformers.loss.loss_utils import nn

nn.functional.cross_entropy = liger_cross_entropy

if fused_linear_cross_entropy:
if model is not None:
model.forward = MethodType(deepseek_v3_lce_forward, model)
else:
modeling_deepseek_v3.DeepseekV3ForCausalLM.forward = deepseek_v3_lce_forward

if swiglu:
modeling_deepseek_v3.DeepseekV3MLP = LigerQwen3MoeSwiGLUMLP

if model is not None:
base_model: DeepseekV3Model = getattr(model, model.base_model_prefix, model)

if rms_norm:
_patch_rms_norm_module(base_model.norm)
for decoder_layer in base_model.layers:
if swiglu:
shared_experts = getattr(decoder_layer.mlp, "shared_experts", None)
if shared_experts is not None:
_patch_swiglu_module(shared_experts, LigerQwen3MoeSwiGLUMLP)
elif not hasattr(decoder_layer.mlp, "experts"):
_patch_swiglu_module(decoder_layer.mlp, LigerQwen3MoeSwiGLUMLP)
if rms_norm:
_patch_rms_norm_module(decoder_layer.input_layernorm)
_patch_rms_norm_module(decoder_layer.post_attention_layernorm)
q_a_layernorm = getattr(decoder_layer.self_attn, "q_a_layernorm", None)
if q_a_layernorm is not None:
_patch_rms_norm_module(q_a_layernorm)
kv_a_layernorm = getattr(decoder_layer.self_attn, "kv_a_layernorm", None)
if kv_a_layernorm is not None:
_patch_rms_norm_module(kv_a_layernorm)


def apply_liger_kernel_to_deepseek_v4(
rope: bool = False,
cross_entropy: bool = False,
Expand Down Expand Up @@ -3532,6 +3616,7 @@ def __init__(self, hidden_size, eps=1e-6, **kwargs):

# Model type corresponds to the keys defined in transformers/models/auto/modeling_auto.py
MODEL_TYPE_TO_APPLY_LIGER_FN = {
"deepseek_v3": apply_liger_kernel_to_deepseek_v3,
"deepseek_v4": apply_liger_kernel_to_deepseek_v4,
"gemma": apply_liger_kernel_to_gemma,
"gemma2": apply_liger_kernel_to_gemma2,
Expand Down
66 changes: 65 additions & 1 deletion test/convergence/bf16/test_mini_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
from transformers.models.qwen2 import Qwen2Config
from transformers.models.qwen2 import Qwen2ForCausalLM

from liger_kernel.transformers import apply_liger_kernel_to_deepseek_v3
from liger_kernel.transformers import apply_liger_kernel_to_deepseek_v4
from liger_kernel.transformers import apply_liger_kernel_to_exaone4
from liger_kernel.transformers import apply_liger_kernel_to_falcon_h1
Expand Down Expand Up @@ -68,6 +69,7 @@
from test.utils import get_logprobs
from test.utils import get_topk
from test.utils import require_deterministic
from test.utils import revert_liger_kernel_to_deepseek_v3
from test.utils import revert_liger_kernel_to_deepseek_v4
from test.utils import revert_liger_kernel_to_exaone4
from test.utils import revert_liger_kernel_to_falcon_h1
Expand Down Expand Up @@ -340,6 +342,14 @@
except ImportError:
HUNYUAN_V1_AVAILABLE = False

try:
from transformers.models.deepseek_v3.configuration_deepseek_v3 import DeepseekV3Config
from transformers.models.deepseek_v3.modeling_deepseek_v3 import DeepseekV3ForCausalLM

DEEPSEEK_V3_AVAILABLE = True
except ImportError:
DEEPSEEK_V3_AVAILABLE = False

try:
from transformers.models.deepseek_v4.configuration_deepseek_v4 import DeepseekV4Config
from transformers.models.deepseek_v4.modeling_deepseek_v4 import DeepseekV4ForCausalLM
Expand Down Expand Up @@ -1648,6 +1658,35 @@
),
)

if DEEPSEEK_V3_AVAILABLE:
MINI_MODEL_SETUPS["mini_deepseek_v3"] = MiniModelConfig(
liger_kernel_patch_func=apply_liger_kernel_to_deepseek_v3,
liger_kernel_patch_revert_func=revert_liger_kernel_to_deepseek_v3,
model_class=DeepseekV3ForCausalLM,
mini_model_config=DeepseekV3Config(
vocab_size=32000,
hidden_size=32,
intermediate_size=64,
moe_intermediate_size=16,
num_hidden_layers=4,
num_attention_heads=2,
num_key_value_heads=2,
q_lora_rank=8,
kv_lora_rank=8,
qk_rope_head_dim=8,
qk_nope_head_dim=8,
v_head_dim=16,
num_experts_per_tok=2,
n_routed_experts=4,
n_shared_experts=1,
n_group=2,
topk_group=1,
first_k_dense_replace=1,
max_position_embeddings=128,
attn_implementation="sdpa",
),
)

if DEEPSEEK_V4_AVAILABLE:
MINI_MODEL_SETUPS["mini_deepseek_v4"] = MiniModelConfig(
liger_kernel_patch_func=apply_liger_kernel_to_deepseek_v4,
Expand Down Expand Up @@ -1767,7 +1806,13 @@ def run_mini_model(
"rms_norm": True,
}

if "glm4" in model_name or "qwen3_next" in model_name or "qwen3_5" in model_name or "deepseek_v4" in model_name:
if (
"glm4" in model_name
or "qwen3_next" in model_name
or "qwen3_5" in model_name
or "deepseek_v3" in model_name
or "deepseek_v4" in model_name
):
kwargs["rope"] = False

model_supports_layer_norm = "qwen2_vl" in model_name
Expand Down Expand Up @@ -2467,6 +2512,25 @@ def run_mini_model(
),
],
),
pytest.param(
"mini_deepseek_v3",
32,
1e-5,
torch.bfloat16,
1e-2,
5e-2,
1e-1,
1e-2,
1e-2,
1e-2,
marks=[
pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"),
pytest.mark.skipif(
not DEEPSEEK_V3_AVAILABLE,
reason="DeepSeek-V3 not available in this version of transformers",
),
],
),
pytest.param(
"mini_deepseek_v4",
32,
Expand Down
Loading