diff --git a/benchmark/scripts/benchmark_megatron_fused_linear_cross_entropy.py b/benchmark/scripts/benchmark_megatron_fused_linear_cross_entropy.py new file mode 100644 index 000000000..b4e846e8f --- /dev/null +++ b/benchmark/scripts/benchmark_megatron_fused_linear_cross_entropy.py @@ -0,0 +1,559 @@ +"""Benchmark hidden-to-loss FLCE against Megatron's materialized output-loss stack. + +Megatron-Core does not provide a fused linear cross-entropy kernel. Its comparable +training path is: + + vocab-parallel linear -> materialized local logits -> fused vocab-parallel CE + +This script compares that path with ``LigerMegatronFusedLinearCrossEntropy``, +which uses portable Triton kernels and saves low-precision CE state to avoid +projection recomputation. ``liger-cutile`` replaces all local compute with +CuTile, while ``liger-cutedsl`` uses a persistent SM100 CuTe DSL projection. +When Megatron-Core is installed, the ``megatron-core`` provider uses its fused +CE. The always-available ``megatron-compatible`` provider uses Liger's drop-in +Megatron CE. + +Backward timing creates a fresh graph outside each timed event pair, so only +backward execution is measured while respecting Megatron's single-use fused CE +graph. Fixed iteration counts keep all tensor-parallel ranks in collective +lockstep. + +Examples: + + python benchmark_megatron_fused_linear_cross_entropy.py --tp-size 1 + torchrun --help # not needed; the script spawns TP ranks itself + python benchmark_megatron_fused_linear_cross_entropy.py --tp-size 4 \ + --token-counts 512 2048 --vocab-sizes 32000 128256 +""" + +from __future__ import annotations + +import argparse +import gc +import os +import tempfile + +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path + +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch.nn.functional as F + +from utils import BenchmarkData +from utils import get_formatted_time +from utils import get_gpu_name +from utils import update_benchmark_data_csv + +from liger_kernel.megatron import LigerMegatronCrossEntropy +from liger_kernel.megatron import LigerMegatronFusedLinearCrossEntropy + +try: + from liger_kernel.ops.cutedsl.ops.megatron_fused_linear_cross_entropy import ( + liger_megatron_fused_linear_cross_entropy as cutedsl_megatron_fused_linear_cross_entropy, + ) + + _CUTEDSL_AVAILABLE = True +except ImportError: + cutedsl_megatron_fused_linear_cross_entropy = None + _CUTEDSL_AVAILABLE = False + +try: + from liger_kernel.ops.cutile.ops.megatron_fused_linear_cross_entropy import ( + liger_megatron_fused_linear_cross_entropy as cutile_megatron_fused_linear_cross_entropy, + ) + + _CUTILE_AVAILABLE = True +except ImportError: + cutile_megatron_fused_linear_cross_entropy = None + _CUTILE_AVAILABLE = False + +try: + from megatron.core.fusions.fused_cross_entropy import fused_vocab_parallel_cross_entropy + + _MEGATRON_CORE_AVAILABLE = True +except ImportError: + fused_vocab_parallel_cross_entropy = None + _MEGATRON_CORE_AVAILABLE = False + + +_SPEED_SAMPLES = 5 +_MEMORY_SAMPLES = 3 +_DTYPES = { + "bf16": torch.bfloat16, + "fp16": torch.float16, +} + + +@dataclass +class _ProviderState: + hidden: torch.Tensor + weight: torch.Tensor + bias: torch.Tensor | None + target: torch.Tensor + forward: Callable[[], torch.Tensor] + + def clear_grads(self) -> None: + self.hidden.grad = None + self.weight.grad = None + if self.bias is not None: + self.bias.grad = None + + +def _all_reduce_hidden_grad(grad: torch.Tensor, tp_group): + dist.all_reduce(grad, op=dist.ReduceOp.SUM, group=tp_group) + return grad + + +def _make_state( + provider: str, + hidden_master: torch.Tensor, + weight_master: torch.Tensor, + bias_master: torch.Tensor | None, + target: torch.Tensor, + tp_group, + tp_size: int, +) -> _ProviderState: + hidden = hidden_master.clone().requires_grad_(True) + weight = weight_master.clone().requires_grad_(True) + bias = bias_master.clone().requires_grad_(True) if bias_master is not None else None + + if provider == "liger": + loss = LigerMegatronFusedLinearCrossEntropy() + forward = lambda: loss(hidden, weight, target, bias=bias, tp_group=tp_group) + elif provider == "liger-cutedsl": + if not _CUTEDSL_AVAILABLE: + raise RuntimeError("provider 'liger-cutedsl' requires nvidia-cutlass-dsl.") + forward = lambda: cutedsl_megatron_fused_linear_cross_entropy( + hidden, + weight, + target, + bias=bias, + tp_group=tp_group, + ) + elif provider == "liger-cutile": + if not _CUTILE_AVAILABLE: + raise RuntimeError("provider 'liger-cutile' requires the cuda-tile package.") + forward = lambda: cutile_megatron_fused_linear_cross_entropy( + hidden, + weight, + target, + bias=bias, + tp_group=tp_group, + ) + else: + if tp_size > 1: + hidden.register_hook(lambda grad: _all_reduce_hidden_grad(grad, tp_group)) + + if provider == "megatron-core": + if not _MEGATRON_CORE_AVAILABLE: + raise RuntimeError("provider 'megatron-core' requires the megatron-core package.") + ce_forward = lambda logits: fused_vocab_parallel_cross_entropy(logits, target, tp_group) + elif provider == "megatron-compatible": + ce = LigerMegatronCrossEntropy() + ce_forward = lambda logits: ce(logits, target, tp_group=tp_group) + else: + raise ValueError(f"unknown provider: {provider!r}") + + forward = lambda: ce_forward(F.linear(hidden, weight, bias)) + + return _ProviderState(hidden, weight, bias, target, forward) + + +def _synchronized_elapsed_ms(step, tp_group, iterations: int) -> float: + dist.barrier(group=tp_group) + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(iterations): + step() + end.record() + torch.cuda.synchronize() + elapsed = torch.tensor(start.elapsed_time(end) / iterations, device="cuda") + dist.all_reduce(elapsed, op=dist.ReduceOp.MAX, group=tp_group) + return float(elapsed) + + +def _synchronized_backward_ms(state: _ProviderState, tp_group, iterations: int) -> float: + dist.barrier(group=tp_group) + torch.cuda.synchronize() + event_pairs = [] + for _ in range(iterations): + state.clear_grads() + loss = state.forward() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + loss.backward(torch.ones_like(loss)) + end.record() + event_pairs.append((start, end)) + torch.cuda.synchronize() + elapsed = torch.tensor( + sum(start.elapsed_time(end) for start, end in event_pairs) / iterations, + device="cuda", + ) + dist.all_reduce(elapsed, op=dist.ReduceOp.MAX, group=tp_group) + return float(elapsed) + + +def _quantiles(samples: list[float]) -> tuple[float, float, float]: + values = torch.tensor(samples) + return ( + float(values.quantile(0.5)), + float(values.quantile(0.2)), + float(values.quantile(0.8)), + ) + + +def _speed( + state: _ProviderState, + tp_group, + warmup_iterations: int, + measure_iterations: int, +) -> dict[str, tuple[float, float, float]]: + def forward_step(): + state.forward() + + def full_step(): + state.clear_grads() + loss = state.forward() + loss.backward(torch.ones_like(loss)) + + for _ in range(warmup_iterations): + full_step() + torch.cuda.synchronize() + + forward_samples = [ + _synchronized_elapsed_ms(forward_step, tp_group, measure_iterations) for _ in range(_SPEED_SAMPLES) + ] + backward_samples = [_synchronized_backward_ms(state, tp_group, measure_iterations) for _ in range(_SPEED_SAMPLES)] + full_samples = [_synchronized_elapsed_ms(full_step, tp_group, measure_iterations) for _ in range(_SPEED_SAMPLES)] + return { + "forward": _quantiles(forward_samples), + "backward": _quantiles(backward_samples), + "full": _quantiles(full_samples), + } + + +def _memory(state: _ProviderState, tp_group) -> tuple[float, float, float]: + def full_step(): + state.clear_grads() + loss = state.forward() + loss.backward(torch.ones_like(loss)) + + full_step() + torch.cuda.synchronize() + state.clear_grads() + gc.collect() + samples = [] + for _ in range(_MEMORY_SAMPLES): + state.clear_grads() + gc.collect() + torch.cuda.reset_peak_memory_stats() + full_step() + torch.cuda.synchronize() + peak = torch.tensor(torch.cuda.max_memory_allocated() / 2**20, device="cuda") + dist.all_reduce(peak, op=dist.ReduceOp.MAX, group=tp_group) + samples.append(float(peak)) + return _quantiles(samples) + + +def _make_masters( + rank: int, + tp_size: int, + num_tokens: int, + hidden_size: int, + vocab_global: int, + dtype: torch.dtype, + with_bias: bool, + device: torch.device, +): + if vocab_global % tp_size: + raise ValueError(f"vocab size {vocab_global} must be divisible by TP={tp_size}.") + vocab_local = vocab_global // tp_size + generator = torch.Generator(device=device) + generator.manual_seed(17) + hidden = torch.randn( + num_tokens, + 1, + hidden_size, + device=device, + dtype=dtype, + generator=generator, + ) + target = torch.randint( + vocab_global, + (num_tokens, 1), + device=device, + dtype=torch.long, + generator=generator, + ) + dist.broadcast(hidden, src=0) + dist.broadcast(target, src=0) + + generator.manual_seed(1000 + rank) + weight = torch.randn( + vocab_local, + hidden_size, + device=device, + dtype=dtype, + generator=generator, + ) + bias = torch.randn(vocab_local, device=device, dtype=dtype, generator=generator) if with_bias else None + return hidden, weight, bias, target + + +def _check_correctness( + rank: int, + tp_size: int, + tp_group, + dtype: torch.dtype, + device: torch.device, + providers, +): + hidden, weight, bias, target = _make_masters( + rank, + tp_size, + num_tokens=32, + hidden_size=256, + vocab_global=1024, + dtype=dtype, + with_bias=True, + device=device, + ) + weight.mul_(0.02) + bias.mul_(0.02) + upstream = torch.randn_like(target, dtype=torch.float32) + dist.broadcast(upstream, src=0) + outputs = {} + correctness_providers = ["megatron-compatible"] + correctness_providers.extend(provider for provider in providers if provider != "megatron-compatible") + for provider in correctness_providers: + state = _make_state(provider, hidden, weight, bias, target, tp_group, tp_size) + loss = state.forward() + loss.backward(upstream) + outputs[provider] = ( + loss.detach().float(), + state.hidden.grad.detach().float(), + state.weight.grad.detach().float(), + state.bias.grad.detach().float(), + ) + + reference = outputs["megatron-compatible"] + names = ("loss", "grad_hidden", "grad_weight", "grad_bias") + for provider in correctness_providers[1:]: + actual = outputs[provider] + for name, actual_tensor, reference_tensor in zip(names, actual, reference): + torch.testing.assert_close( + actual_tensor, + reference_tensor, + atol=5e-3, + rtol=5e-2, + msg=f"{provider}: {name}", + ) + + +def _worker( + rank, + tp_size, + providers, + token_counts, + vocab_sizes, + hidden_size, + dtype_name, + with_bias, + warmup_iterations, + measure_iterations, + rendezvous, + result_path, + overwrite, +): + os.environ.setdefault("MASTER_ADDR", "localhost") + os.environ.setdefault("MASTER_PORT", "29500") + dist.init_process_group( + backend="nccl", + init_method=f"file://{rendezvous}", + rank=rank, + world_size=tp_size, + ) + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + tp_group = dist.group.WORLD + dtype = _DTYPES[dtype_name] + + _check_correctness(rank, tp_size, tp_group, dtype, device, providers) + if rank == 0: + print("Correctness: loss and gradients match the materialized reference.", flush=True) + + grouped_speed = {} if rank == 0 else None + grouped_memory = {} if rank == 0 else None + for num_tokens in token_counts: + for vocab_global in vocab_sizes: + masters = _make_masters( + rank, + tp_size, + num_tokens, + hidden_size, + vocab_global, + dtype, + with_bias, + device, + ) + for provider in providers: + state = _make_state(provider, *masters, tp_group, tp_size) + speed = _speed(state, tp_group, warmup_iterations, measure_iterations) + memory = _memory(state, tp_group) + if rank == 0: + for mode, (p50, p20, p80) in speed.items(): + grouped_speed.setdefault((provider, mode, num_tokens), []).append((vocab_global, p50, p20, p80)) + print( + f"[speed] {provider:>21s} TP={tp_size} BT={num_tokens:>5d} " + f"V={vocab_global:>6d} {mode:>8s}: {p50:.4f} ms", + flush=True, + ) + grouped_memory.setdefault((provider, "full", num_tokens), []).append((vocab_global, *memory)) + print( + f"[memory] {provider:>21s} TP={tp_size} BT={num_tokens:>5d} " + f"V={vocab_global:>6d}: {memory[0]:.1f} MB", + flush=True, + ) + del state + torch.cuda.empty_cache() + dist.barrier(group=tp_group) + del masters + + if rank == 0: + timestamp = get_formatted_time() + gpu_name = get_gpu_name() + rows = [] + common = { + "kernel_name": "megatron_fused_linear_cross_entropy", + "gpu_name": gpu_name, + "x_name": "V", + "x_label": "global vocab size", + "timestamp": timestamp, + } + for (provider, mode, num_tokens), samples in grouped_speed.items(): + samples.sort() + rows.append( + BenchmarkData( + kernel_provider=provider, + metric_name="speed", + metric_unit="ms", + x_values=[row[0] for row in samples], + y_values_50=[row[1] for row in samples], + y_values_20=[row[2] for row in samples], + y_values_80=[row[3] for row in samples], + kernel_operation_mode=mode, + extra_benchmark_config_str=( + f'{{"BT": {num_tokens}, "H": {hidden_size}, "TP": {tp_size}, ' + f'"dtype": "{dtype_name}", "bias": {str(with_bias).lower()}}}' + ), + **common, + ) + ) + for (provider, mode, num_tokens), samples in grouped_memory.items(): + samples.sort() + rows.append( + BenchmarkData( + kernel_provider=provider, + metric_name="memory", + metric_unit="MB", + x_values=[row[0] for row in samples], + y_values_50=[row[1] for row in samples], + y_values_20=[row[2] for row in samples], + y_values_80=[row[3] for row in samples], + kernel_operation_mode=mode, + extra_benchmark_config_str=( + f'{{"BT": {num_tokens}, "H": {hidden_size}, "TP": {tp_size}, ' + f'"dtype": "{dtype_name}", "bias": {str(with_bias).lower()}}}' + ), + **common, + ) + ) + update_benchmark_data_csv(rows, filename=result_path, overwrite=overwrite) + + dist.destroy_process_group() + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--tp-size", type=int, default=1) + parser.add_argument("--token-counts", type=int, nargs="+", default=[512, 2048]) + parser.add_argument("--hidden-size", type=int, default=4096) + parser.add_argument("--vocab-sizes", type=int, nargs="+", default=[32000, 128256]) + parser.add_argument("--dtype", choices=sorted(_DTYPES), default="bf16") + parser.add_argument("--with-bias", action="store_true") + parser.add_argument("--warmup-iterations", type=int, default=3) + parser.add_argument("--measure-iterations", type=int, default=10) + parser.add_argument( + "--providers", + nargs="+", + choices=[ + "liger", + "liger-cutedsl", + "liger-cutile", + "megatron-compatible", + "megatron-core", + ], + ) + parser.add_argument( + "--output", + type=Path, + default=Path(__file__).resolve().parents[1] / "data" / "all_benchmark_data_megatron_flce.csv", + ) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + + if args.tp_size > torch.cuda.device_count(): + raise RuntimeError(f"--tp-size={args.tp_size} requires {args.tp_size} GPUs; found {torch.cuda.device_count()}.") + providers = args.providers or ["megatron-compatible", "liger"] + if _MEGATRON_CORE_AVAILABLE and args.providers is None: + providers.insert(1, "megatron-core") + if "megatron-core" in providers and not _MEGATRON_CORE_AVAILABLE: + raise RuntimeError("provider 'megatron-core' requested, but megatron-core is not installed.") + if "liger-cutedsl" in providers and not _CUTEDSL_AVAILABLE: + raise RuntimeError("provider 'liger-cutedsl' requested, but nvidia-cutlass-dsl is not installed.") + if "liger-cutedsl" in providers: + unsupported_devices = [ + index for index in range(args.tp_size) if torch.cuda.get_device_capability(index) != (10, 0) + ] + if unsupported_devices: + raise RuntimeError( + "provider 'liger-cutedsl' requires SM100 (compute capability 10.0); " + f"unsupported CUDA device indices: {unsupported_devices}." + ) + if "liger-cutile" in providers and not _CUTILE_AVAILABLE: + raise RuntimeError("provider 'liger-cutile' requested, but cuda-tile is not installed.") + if min(args.token_counts) <= 0 or args.hidden_size <= 0 or min(args.vocab_sizes) <= 0: + raise ValueError("token counts, hidden size, and vocabulary sizes must be positive.") + + output = args.output.resolve() + output.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile() as rendezvous: + mp.spawn( + _worker, + args=( + args.tp_size, + providers, + args.token_counts, + args.vocab_sizes, + args.hidden_size, + args.dtype, + args.with_bias, + args.warmup_iterations, + args.measure_iterations, + rendezvous.name, + str(output), + args.overwrite, + ), + nprocs=args.tp_size, + join=True, + ) + + +if __name__ == "__main__": + main() diff --git a/docs/High-Level-APIs.md b/docs/High-Level-APIs.md index 6bbe008a9..64029f60b 100644 --- a/docs/High-Level-APIs.md +++ b/docs/High-Level-APIs.md @@ -98,38 +98,57 @@ You can also use the Patching APIs to use the kernels for a specific model archi Liger also exposes a patch for the [Megatron-LM](https://github.com/NVIDIA/Megatron-LM) training framework, replacing Megatron's native RMSNorm and both vocab-parallel -cross-entropy paths (fused and unfused) with Liger's Triton kernels. +cross-entropy paths (fused and unfused) with Liger kernels. It can also fuse +the GPT output projection with cross-entropy. | **Framework** | **API** | **Supported Operations** | |---------------|--------------------------------------------------------|--------------------------| -| Megatron-LM | `liger_kernel.megatron.apply_liger_kernel_to_megatron` | RMSNorm, CrossEntropyLoss | +| Megatron-LM | `liger_kernel.megatron.apply_liger_kernel_to_megatron` | RMSNorm, CrossEntropyLoss, fused output projection + CrossEntropyLoss | +| Megatron-LM | `liger_kernel.megatron.LigerMegatronFusedLinearCrossEntropy` | Fused output projection + CrossEntropyLoss | -**Scope**: Initial release supports `tensor_model_parallel_size=1` only for -cross-entropy. Vocab-parallel cross-entropy (TP>1) is follow-up work — with -TP>1, each rank holds a sharded `[N, V/tp]` logits slice and cross-entropy -requires cross-rank all-reduces that Liger's kernel does not perform. The -patch raises a `RuntimeError` at patch time or call time if TP>1 is detected. +Both cross-entropy patches and FLCE support TP1 and TP>1. Automatic FLCE +patching requires Megatron-Core 0.18 or newer, BF16 or FP16, a native +`ColumnParallelLinear` output layer, and no sequence parallelism, +gradient-accumulation fusion, or gathered logits. Unsupported configurations +raise an explicit error. **Usage**: ```python from liger_kernel.megatron import apply_liger_kernel_to_megatron -# Call before Megatron's forward pass reaches compute_language_model_loss. -# Defaults match Megatron's native CE behavior; no CE-specific config needed. -apply_liger_kernel_to_megatron(rms_norm=True, cross_entropy=True) +apply_liger_kernel_to_megatron( + rms_norm=True, + cross_entropy=True, + fused_linear_cross_entropy=True, +) ``` Both the fused (`config.cross_entropy_loss_fusion=True`, `cross_entropy_fusion_impl='native'`) and unfused (`config.cross_entropy_loss_fusion=False`) CE paths are patched in a single -call, so Megatron picks up Liger regardless of which path your config selects. +call. When `fused_linear_cross_entropy=True`, labeled standard GPT forwards +bypass both paths and use FLCE directly. For custom integrations, use +`LigerMegatronFusedLinearCrossEntropy` with replicated hidden states, a +contiguous local vocabulary shard, and global target indices. For training setups that need explicit kernel configuration (custom `ignore_index`, `label_smoothing`, etc.), instantiate `LigerMegatronCrossEntropy` directly and wire it into your model — see `examples/megatron/run_mode2_hand_spec.py`. +::: liger_kernel.megatron.LigerMegatronFusedLinearCrossEntropy + options: + extra: + show_docstring: true + show_signature: true + +::: liger_kernel.megatron.liger_megatron_fused_linear_cross_entropy_output_processor + options: + extra: + show_docstring: true + show_signature: true + ::: liger_kernel.megatron.apply_liger_kernel_to_megatron options: extra: diff --git a/src/liger_kernel/megatron/__init__.py b/src/liger_kernel/megatron/__init__.py index 94f3da552..e54fbb2ab 100644 --- a/src/liger_kernel/megatron/__init__.py +++ b/src/liger_kernel/megatron/__init__.py @@ -10,10 +10,14 @@ experts. Mode 1 patches ``fused_bias_swiglu.SwiGLUFunction`` instead; both fall back to Megatron for FP8 input store and CPU activation offload, and neither touches the bias or MoE-routed variants. + LigerMegatronFusedLinearCrossEntropy — hidden-state-to-loss fused output + projection for tensor-parallel vocabulary shards. + liger_megatron_fused_linear_cross_entropy_output_processor — adapts FLCE + to Megatron's GPT output-processor hook. apply_liger_kernel_to_megatron — patches Megatron-Core so existing training scripts pick up Liger kernels with one line. Currently supports RMSNorm (via BackendSpecProvider), both the fused and unfused - vocab-parallel cross-entropy paths, and SwiGLU. + vocab-parallel cross-entropy paths, opt-in GPT FLCE, and SwiGLU. The general-purpose ``LigerVocabParallelCrossEntropy`` Module lives under ``liger_kernel.transformers`` alongside the other nn.Module wrappers; the @@ -23,13 +27,17 @@ """ from liger_kernel.megatron.cross_entropy import LigerMegatronCrossEntropy +from liger_kernel.megatron.fused_linear_cross_entropy import LigerMegatronFusedLinearCrossEntropy +from liger_kernel.megatron.fused_linear_cross_entropy import liger_megatron_fused_linear_cross_entropy_output_processor from liger_kernel.megatron.monkey_patch import apply_liger_kernel_to_megatron from liger_kernel.megatron.rms_norm import LigerMegatronRMSNorm from liger_kernel.megatron.swiglu import LigerMegatronSwiGLU __all__ = [ "LigerMegatronCrossEntropy", + "LigerMegatronFusedLinearCrossEntropy", "LigerMegatronRMSNorm", "LigerMegatronSwiGLU", "apply_liger_kernel_to_megatron", + "liger_megatron_fused_linear_cross_entropy_output_processor", ] diff --git a/src/liger_kernel/megatron/fused_linear_cross_entropy.py b/src/liger_kernel/megatron/fused_linear_cross_entropy.py new file mode 100644 index 000000000..0e5507344 --- /dev/null +++ b/src/liger_kernel/megatron/fused_linear_cross_entropy.py @@ -0,0 +1,100 @@ +"""Megatron-facing module for tensor-parallel fused linear cross entropy.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + +from liger_kernel.ops import LigerMegatronFusedLinearCrossEntropyFunction + + +class LigerMegatronFusedLinearCrossEntropy(nn.Module): + """Fuse a vocab-sharded output projection with per-token cross entropy. + + ``hidden`` is replicated across TP ranks and ``weight`` contains the local + contiguous vocabulary shard. The output shape matches ``target``. + """ + + def __init__( + self, + ignore_index: int = -100, + ): + super().__init__() + self.ignore_index = ignore_index + + def forward( + self, + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None = None, + tp_group=None, + ) -> torch.Tensor: + return LigerMegatronFusedLinearCrossEntropyFunction.apply( + hidden, + weight, + target, + bias, + tp_group, + self.ignore_index, + ) + + def extra_repr(self) -> str: + return f"ignore_index={self.ignore_index}" + + +def liger_megatron_fused_linear_cross_entropy_output_processor( + *, + hidden_states: torch.Tensor, + output_layer, + output_weight: torch.Tensor | None, + labels: torch.Tensor, + runtime_gather_output: bool | None, + config, + **_, +) -> torch.Tensor: + """Megatron ``GPTModel`` output processor for the native TP output layer.""" + unsupported = [] + if type(output_layer).__name__ != "ColumnParallelLinear": + unsupported.append("the output layer is not Megatron's native ColumnParallelLinear") + if getattr(output_layer, "sequence_parallel", False): + unsupported.append("sequence_parallel=True") + if getattr(output_layer, "gradient_accumulation_fusion", False): + unsupported.append("gradient_accumulation_fusion=True") + if getattr(output_layer, "disable_grad_reduce", False): + unsupported.append("output-layer dgrad reduction is disabled") + if getattr(output_layer, "explicit_expert_comm", False): + unsupported.append("the output layer uses explicit expert communication") + if getattr(output_layer, "skip_bias_add", False): + unsupported.append("the output layer returns bias separately") + if getattr(config, "defer_embedding_wgrad_compute", False): + unsupported.append("defer_embedding_wgrad_compute=True") + if getattr(config, "mtp_num_layers", None): + unsupported.append("MTP is enabled") + if getattr(config, "use_mup", False): + unsupported.append("MuP output scaling is enabled") + + gather_output = ( + getattr(output_layer, "gather_output", False) if runtime_gather_output is None else runtime_gather_output + ) + if gather_output: + unsupported.append("the output layer gathers TP logits") + if unsupported: + raise RuntimeError( + "Liger Megatron FLCE does not support this GPT output configuration: " + "; ".join(unsupported) + ) + + weight = output_weight if output_weight is not None else getattr(output_layer, "weight", None) + if weight is None: + raise RuntimeError("Liger Megatron FLCE requires an output weight tensor.") + + labels_sb = labels.transpose(0, 1).contiguous() + loss_sb = LigerMegatronFusedLinearCrossEntropyFunction.apply( + hidden_states, + weight, + labels_sb, + getattr(output_layer, "bias", None), + getattr(output_layer, "tp_group", None), + -100, + ) + return loss_sb.transpose(0, 1).contiguous() diff --git a/src/liger_kernel/megatron/monkey_patch.py b/src/liger_kernel/megatron/monkey_patch.py index e97567ca3..7b4db807c 100644 --- a/src/liger_kernel/megatron/monkey_patch.py +++ b/src/liger_kernel/megatron/monkey_patch.py @@ -2,16 +2,27 @@ from __future__ import annotations +import functools +import inspect import logging +import sys logger = logging.getLogger(__name__) _PATCH_MARKER = "__liger_patched__" +def _replace_loaded_binding(module_name: str, symbol_name: str, original, replacement) -> None: + """Update a known by-name import without clobbering third-party patches.""" + module = sys.modules.get(module_name) + if module is not None and getattr(module, symbol_name, None) is original: + setattr(module, symbol_name, replacement) + + def apply_liger_kernel_to_megatron( rms_norm: bool = True, cross_entropy: bool = False, + fused_linear_cross_entropy: bool = False, swiglu: bool = False, ) -> None: """Patch Megatron-Core to use Liger Triton kernels. @@ -39,6 +50,11 @@ def apply_liger_kernel_to_megatron( wrapper additionally honors a runtime ``label_smoothing`` argument, matching native's ``(logits, target, label_smoothing=0.0, tp_group=None)``. + fused_linear_cross_entropy: When ``True`` inject Liger's fused local + output projection and vocab-parallel cross-entropy through + ``GPTModel._postprocess`` for supported native + ``ColumnParallelLinear`` training configurations. This is opt-in + and independent of ``cross_entropy``. swiglu: When ``True`` replace ``megatron.core.fusions.fused_bias_swiglu.SwiGLUFunction`` with Liger's Triton SiLU-multiply kernel, covering the dense ``MLP`` and the MoE @@ -68,6 +84,8 @@ def apply_liger_kernel_to_megatron( if cross_entropy: _patch_fused_vocab_parallel_cross_entropy() _patch_vocab_parallel_cross_entropy() + if fused_linear_cross_entropy: + _patch_gpt_fused_linear_cross_entropy() if swiglu: _patch_swiglu_function() @@ -183,7 +201,14 @@ def _patch_fused_vocab_parallel_cross_entropy() -> None: ) if getattr(fused_ce.fused_vocab_parallel_cross_entropy, _PATCH_MARKER, False): - return # already patched + replacement = fused_ce.fused_vocab_parallel_cross_entropy + _replace_loaded_binding( + "megatron.core.models.common.language_module.language_module", + "fused_vocab_parallel_cross_entropy", + replacement.__wrapped__, + replacement, + ) + return original = fused_ce.fused_vocab_parallel_cross_entropy @@ -197,11 +222,15 @@ def liger_fused_vocab_parallel_cross_entropy(vocab_parallel_logits, target, tp_g setattr(liger_fused_vocab_parallel_cross_entropy, _PATCH_MARKER, True) setattr(liger_fused_vocab_parallel_cross_entropy, "__wrapped__", original) fused_ce.fused_vocab_parallel_cross_entropy = liger_fused_vocab_parallel_cross_entropy - - logger.info( - "Patched megatron.core.fusions.fused_cross_entropy.fused_vocab_parallel_cross_entropy with Liger cross-entropy." + _replace_loaded_binding( + "megatron.core.models.common.language_module.language_module", + "fused_vocab_parallel_cross_entropy", + original, + liger_fused_vocab_parallel_cross_entropy, ) + logger.info("Patched Megatron's fused vocab-parallel cross-entropy definition and loaded LanguageModule binding.") + def _patch_vocab_parallel_cross_entropy() -> None: """Replace ``megatron.core.tensor_parallel.cross_entropy.vocab_parallel_cross_entropy``. @@ -229,7 +258,14 @@ def _patch_vocab_parallel_cross_entropy() -> None: ) if getattr(unfused_ce.vocab_parallel_cross_entropy, _PATCH_MARKER, False): - return # already patched + replacement = unfused_ce.vocab_parallel_cross_entropy + _replace_loaded_binding( + "megatron.core.tensor_parallel", + "vocab_parallel_cross_entropy", + replacement.__wrapped__, + replacement, + ) + return original = unfused_ce.vocab_parallel_cross_entropy @@ -260,11 +296,58 @@ def liger_vocab_parallel_cross_entropy( setattr(liger_vocab_parallel_cross_entropy, _PATCH_MARKER, True) setattr(liger_vocab_parallel_cross_entropy, "__wrapped__", original) unfused_ce.vocab_parallel_cross_entropy = liger_vocab_parallel_cross_entropy + _replace_loaded_binding( + "megatron.core.tensor_parallel", + "vocab_parallel_cross_entropy", + original, + liger_vocab_parallel_cross_entropy, + ) - logger.info( - "Patched megatron.core.tensor_parallel.cross_entropy.vocab_parallel_cross_entropy with Liger cross-entropy." + logger.info("Patched Megatron's unfused vocab-parallel cross-entropy definition and tensor_parallel export.") + + +def _patch_gpt_fused_linear_cross_entropy() -> None: + """Inject Liger FLCE through Megatron's GPT output-processor hook.""" + try: + import megatron.core.models.gpt.gpt_model as gpt_model + except ImportError as exc: + raise ImportError( + "apply_liger_kernel_to_megatron(fused_linear_cross_entropy=True) requires " + "megatron-core with megatron.core.models.gpt.gpt_model.GPTModel." + ) from exc + + if not hasattr(gpt_model, "GPTModel") or not hasattr(gpt_model.GPTModel, "_postprocess"): + raise ImportError( + "megatron.core.models.gpt.gpt_model.GPTModel._postprocess was not found. " + "The symbol path may have changed in your Megatron-Core version." + ) + + current = gpt_model.GPTModel._postprocess + if getattr(current, _PATCH_MARKER, False): + return + signature = inspect.signature(current) + if not {"labels", "output_processor"} <= signature.parameters.keys(): + raise ImportError( + "Megatron GPTModel._postprocess does not expose the labels and output_processor hooks " + "required by Liger FLCE. Upgrade to Megatron-Core 0.18 or newer." + ) + + from liger_kernel.megatron.fused_linear_cross_entropy import ( + liger_megatron_fused_linear_cross_entropy_output_processor, ) + @functools.wraps(current) + def liger_postprocess(self, *args, **kwargs): + bound = signature.bind(self, *args, **kwargs) + bound.apply_defaults() + if bound.arguments["labels"] is not None and bound.arguments["output_processor"] is None: + bound.arguments["output_processor"] = liger_megatron_fused_linear_cross_entropy_output_processor + return current(*bound.args, **bound.kwargs) + + setattr(liger_postprocess, _PATCH_MARKER, True) + gpt_model.GPTModel._postprocess = liger_postprocess + logger.info("Patched Megatron GPTModel._postprocess to inject Liger fused linear cross-entropy.") + def _patch_swiglu_function() -> None: """Replace ``megatron.core.fusions.fused_bias_swiglu.SwiGLUFunction`` with Liger. diff --git a/src/liger_kernel/ops/__init__.py b/src/liger_kernel/ops/__init__.py index 6bed68643..ee78f04e0 100644 --- a/src/liger_kernel/ops/__init__.py +++ b/src/liger_kernel/ops/__init__.py @@ -67,6 +67,10 @@ 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.llama4_rope import LigerLlama4RopeFunction # noqa: F401 +from liger_kernel.ops.megatron_fused_linear_cross_entropy import ( + LigerMegatronFusedLinearCrossEntropyFunction, # noqa: F401 +) +from liger_kernel.ops.megatron_fused_linear_cross_entropy import liger_megatron_fused_linear_cross_entropy # noqa: F401 from liger_kernel.ops.mhc import LigerMHCCoeffsFunction # noqa: F401 from liger_kernel.ops.mhc import LigerMHCPostResFunction # noqa: F401 from liger_kernel.ops.mhc import LigerMHCPreFunction # noqa: F401 diff --git a/src/liger_kernel/ops/cutedsl/ops/__init__.py b/src/liger_kernel/ops/cutedsl/ops/__init__.py index a5b051b1a..de25f51bc 100644 --- a/src/liger_kernel/ops/cutedsl/ops/__init__.py +++ b/src/liger_kernel/ops/cutedsl/ops/__init__.py @@ -18,6 +18,10 @@ from liger_kernel.ops.cutedsl.ops.fused_scaled_cross_entropy_sm90 import LigerFusedScaledCrossEntropySM90Function from liger_kernel.ops.cutedsl.ops.fused_scaled_cross_entropy_sm90 import fused_scaled_cross_entropy_backward from liger_kernel.ops.cutedsl.ops.fused_scaled_cross_entropy_sm90 import fused_scaled_cross_entropy_forward +from liger_kernel.ops.cutedsl.ops.megatron_fused_linear_cross_entropy import ( + LigerMegatronFusedLinearCrossEntropyFunction, +) +from liger_kernel.ops.cutedsl.ops.megatron_fused_linear_cross_entropy import liger_megatron_fused_linear_cross_entropy from liger_kernel.ops.cutedsl.ops.rms_norm import LigerRMSNormFunction from liger_kernel.ops.cutedsl.ops.rms_norm import rms_norm_backward from liger_kernel.ops.cutedsl.ops.rms_norm import rms_norm_forward @@ -41,6 +45,8 @@ "LigerFusedScaledCrossEntropySM90Function", "fused_scaled_cross_entropy_backward", "fused_scaled_cross_entropy_forward", + "LigerMegatronFusedLinearCrossEntropyFunction", + "liger_megatron_fused_linear_cross_entropy", "LigerRMSNormFunction", "rms_norm_backward", "rms_norm_forward", diff --git a/src/liger_kernel/ops/cutedsl/ops/megatron_fused_linear_cross_entropy.py b/src/liger_kernel/ops/cutedsl/ops/megatron_fused_linear_cross_entropy.py new file mode 100644 index 000000000..1f0f8fdee --- /dev/null +++ b/src/liger_kernel/ops/cutedsl/ops/megatron_fused_linear_cross_entropy.py @@ -0,0 +1,189 @@ +"""CuTe DSL tensor-parallel fused linear cross entropy for Megatron. + +The SM100 path uses a persistent CuTe DSL GEMM for the local vocabulary +projection, Triton for vocabulary-parallel cross entropy, and NCCL for +tensor-parallel collectives. Backward converts the saved projection buffer to +dlogits in-place. +""" + +from __future__ import annotations + +import cutlass +import cutlass.cute as cute +import torch +import torch.distributed as dist +import torch.nn.functional as F + +from liger_kernel.ops.cutedsl.ops._sm100_gemm import K_ALIGNMENT +from liger_kernel.ops.cutedsl.ops._sm100_gemm import run_epilogue_gemm +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _ce_backward_from_logits +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _ce_forward_stats +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _tp_rank_and_world +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _validate_megatron_flce_inputs + + +@cute.jit +def _identity_epilogue(accumulator, output): + output_dtype = output.element_type + for element in cutlass.range_constexpr(cute.size(accumulator)): + output[element] = accumulator[element].to(output_dtype) + + +def _native_cutedsl_supported(hidden: torch.Tensor, weight: torch.Tensor) -> bool: + if hidden.device.type != "cuda" or hidden.dtype not in (torch.bfloat16, torch.float16): + return False + if weight.device != hidden.device or weight.dtype != hidden.dtype: + return False + try: + return torch.cuda.get_device_capability(hidden.device) == (10, 0) + except (AssertionError, RuntimeError): + return False + + +def _cutedsl_projection( + hidden: torch.Tensor, + weight: torch.Tensor, +) -> torch.Tensor: + padding = (-hidden.shape[1]) % K_ALIGNMENT + if padding: + hidden = F.pad(hidden, (0, padding)) + weight = F.pad(weight, (0, padding)) + logits = torch.empty( + hidden.shape[0], + weight.shape[0], + device=hidden.device, + dtype=hidden.dtype, + ) + run_epilogue_gemm(hidden, weight, logits, _identity_epilogue) + return logits + + +def _materialized_backward(ctx, grad_output: torch.Tensor): + hidden, weight, logits, logits_max, sum_exp, target = ctx.saved_tensors + grad_output_1d = grad_output.contiguous().reshape(-1).float() + _ce_backward_from_logits( + logits, + logits_max, + sum_exp, + target, + grad_output_1d, + ctx.vocab_start, + ctx.ignore_index, + ctx.ce_block_size, + ) + + grad_hidden = torch.mm(logits, weight) + reduce_work = ( + dist.all_reduce( + grad_hidden, + op=dist.ReduceOp.SUM, + group=ctx.tp_group, + async_op=True, + ) + if ctx.tp_world > 1 + else None + ) + grad_weight = logits.t() @ hidden + grad_bias = logits.sum(dim=0, dtype=torch.float32).to(ctx.bias_dtype) if ctx.has_bias else None + if reduce_work is not None: + reduce_work.wait() + + grad_hidden = grad_hidden.reshape(ctx.original_hidden_shape) + return grad_hidden, grad_weight, grad_bias + + +class LigerMegatronFusedLinearCrossEntropyFunction(torch.autograd.Function): + """Megatron FLCE using a persistent CuTe DSL SM100 projection.""" + + @staticmethod + def forward( + ctx, + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None, + tp_group, + ignore_index: int, + ) -> torch.Tensor: + _validate_megatron_flce_inputs(hidden, weight, target, bias) + if not _native_cutedsl_supported(hidden, weight): + raise RuntimeError("CuTe DSL Megatron FLCE requires an SM100 GPU and float16 or bfloat16 inputs.") + + tp_rank, tp_world = _tp_rank_and_world(tp_group) + vocab_local = weight.shape[0] + vocab_global = vocab_local * tp_world + vocab_start = tp_rank * vocab_local + flat_target = target.reshape(-1).contiguous() + valid = flat_target != ignore_index + invalid = valid & ((flat_target < 0) | (flat_target >= vocab_global)) + valid_targets = ~torch.any(invalid) + if hasattr(torch, "_assert_async"): + torch._assert_async(valid_targets, f"non-ignored targets must be in [0, {vocab_global}).") + elif not valid_targets.item(): + raise ValueError(f"non-ignored targets must be in [0, {vocab_global}).") + + original_hidden_shape = hidden.shape + hidden_2d = hidden.reshape(-1, hidden.shape[-1]).contiguous() + weight_2d = weight.contiguous() + bias_1d = bias.contiguous() if bias is not None else None + logits = _cutedsl_projection(hidden_2d, weight_2d) + if bias_1d is not None: + logits.add_(bias_1d) + + logits_max = logits.amax(dim=-1).float() + if tp_world > 1: + dist.all_reduce(logits_max, op=dist.ReduceOp.MAX, group=tp_group) + + from liger_kernel.ops.vocab_parallel_cross_entropy import _select_block_size + + ce_block_size = _select_block_size(vocab_local) + stats = _ce_forward_stats( + logits, + logits_max, + flat_target, + vocab_start, + ignore_index, + ce_block_size, + ) + predicted_logit = stats[0] + sum_exp = stats[1] + if tp_world > 1: + dist.all_reduce(stats, op=dist.ReduceOp.SUM, group=tp_group) + + loss = torch.log(sum_exp) - predicted_logit + loss = torch.where(valid, loss, torch.zeros_like(loss)) + ctx.save_for_backward(hidden_2d, weight_2d, logits, logits_max, sum_exp, flat_target) + ctx.has_bias = bias is not None + ctx.bias_dtype = bias.dtype if bias is not None else None + ctx.tp_group = tp_group + ctx.tp_world = tp_world + ctx.vocab_start = vocab_start + ctx.ignore_index = ignore_index + ctx.ce_block_size = ce_block_size + ctx.original_hidden_shape = original_hidden_shape + ctx.hidden_dtype = hidden.dtype + return loss.reshape(target.shape) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + grad_hidden, grad_weight, grad_bias = _materialized_backward(ctx, grad_output) + return grad_hidden, grad_weight, None, grad_bias, None, None + + +def liger_megatron_fused_linear_cross_entropy( + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None = None, + tp_group=None, + ignore_index: int = -100, +) -> torch.Tensor: + """Compute Megatron FLCE with a CuTe DSL projection and NCCL TP collectives.""" + return LigerMegatronFusedLinearCrossEntropyFunction.apply( + hidden, + weight, + target, + bias, + tp_group, + ignore_index, + ) diff --git a/src/liger_kernel/ops/cutile/ops/__init__.py b/src/liger_kernel/ops/cutile/ops/__init__.py index fcbd1cc84..492a7f97e 100644 --- a/src/liger_kernel/ops/cutile/ops/__init__.py +++ b/src/liger_kernel/ops/cutile/ops/__init__.py @@ -33,6 +33,8 @@ from liger_kernel.ops.cutile.ops.layer_norm import layer_norm_backward from liger_kernel.ops.cutile.ops.layer_norm import layer_norm_forward from liger_kernel.ops.cutile.ops.llama4_rope import LigerLlama4RopeFunction +from liger_kernel.ops.cutile.ops.megatron_fused_linear_cross_entropy import LigerMegatronFusedLinearCrossEntropyFunction +from liger_kernel.ops.cutile.ops.megatron_fused_linear_cross_entropy import liger_megatron_fused_linear_cross_entropy from liger_kernel.ops.cutile.ops.multi_token_attention import LigerMultiTokenAttentionFunction from liger_kernel.ops.cutile.ops.qwen2vl_mrope import LigerQwen2VLMRopeFunction from liger_kernel.ops.cutile.ops.rope import LigerRopeFunction @@ -64,6 +66,8 @@ "layer_norm_backward", "layer_norm_forward", "LigerLlama4RopeFunction", + "LigerMegatronFusedLinearCrossEntropyFunction", + "liger_megatron_fused_linear_cross_entropy", "LigerMultiTokenAttentionFunction", "LigerQwen2VLMRopeFunction", "LigerRopeFunction", diff --git a/src/liger_kernel/ops/cutile/ops/megatron_fused_linear_cross_entropy.py b/src/liger_kernel/ops/cutile/ops/megatron_fused_linear_cross_entropy.py new file mode 100644 index 000000000..a31f26a82 --- /dev/null +++ b/src/liger_kernel/ops/cutile/ops/megatron_fused_linear_cross_entropy.py @@ -0,0 +1,579 @@ +"""CuTile tensor-parallel fused linear cross entropy for Megatron.""" + +from __future__ import annotations + +import math + +import cuda.tile as ct +import torch +import torch.distributed as dist + +from liger_kernel.ops.cutile.ops.utils import _next_power_of_2 +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _tp_rank_and_world +from liger_kernel.ops.megatron_fused_linear_cross_entropy import _validate_megatron_flce_inputs + +ConstBool = ct.Constant[bool] +ConstInt = ct.Constant[int] +LOG2E = 1.4426950408889634 +MAX_ROW_BLOCK_SIZE = 4096 + + +@ct.function +def _matmul_body( + a, + b, + bias, + output, + TILE_M: ConstInt, + TILE_N: ConstInt, + TILE_K: ConstInt, + HAS_BIAS: ConstBool, + SWIZZLE: ConstBool, +): + num_m_tiles = ct.num_tiles(a, axis=0, shape=(TILE_M, TILE_K)) + num_n_tiles = ct.num_tiles(b, axis=1, shape=(TILE_K, TILE_N)) + num_k_tiles = ct.num_tiles(a, axis=1, shape=(TILE_M, TILE_K)) + block = ct.bid(0) + if SWIZZLE: + group_size_m = 8 + blocks_per_group = group_size_m * num_n_tiles + group = block // blocks_per_group + first_tile_m = group * group_size_m + active_group_size_m = min(num_m_tiles - first_tile_m, group_size_m) + tile_m = first_tile_m + (block % active_group_size_m) + tile_n = (block % blocks_per_group) // active_group_size_m + else: + tile_m = block // num_n_tiles + tile_n = block % num_n_tiles + + accumulator = ct.full((TILE_M, TILE_N), 0.0, dtype=ct.float32) + for tile_k in range(num_k_tiles): + a_tile = ct.load( + a, + index=(tile_m, tile_k), + shape=(TILE_M, TILE_K), + padding_mode=ct.PaddingMode.ZERO, + ) + b_tile = ct.load( + b, + index=(tile_k, tile_n), + shape=(TILE_K, TILE_N), + padding_mode=ct.PaddingMode.ZERO, + ) + accumulator = ct.mma(a_tile, b_tile, accumulator) + + if HAS_BIAS: + bias_tile = ct.load( + bias, + index=(tile_n,), + shape=(TILE_N,), + padding_mode=ct.PaddingMode.ZERO, + ) + accumulator = accumulator + ct.astype(bias_tile, ct.float32) + + ct.store(output, index=(tile_m, tile_n), tile=ct.astype(accumulator, output.dtype)) + + +@ct.kernel(num_ctas=1) +def _matmul_1cta_kernel( + a, + b, + bias, + output, + TILE_M: ConstInt, + TILE_N: ConstInt, + TILE_K: ConstInt, + HAS_BIAS: ConstBool, + SWIZZLE: ConstBool, +): + _matmul_body(a, b, bias, output, TILE_M, TILE_N, TILE_K, HAS_BIAS, SWIZZLE) + + +@ct.kernel(num_ctas=2) +def _matmul_2cta_kernel( + a, + b, + bias, + output, + TILE_M: ConstInt, + TILE_N: ConstInt, + TILE_K: ConstInt, + HAS_BIAS: ConstBool, + SWIZZLE: ConstBool, +): + _matmul_body(a, b, bias, output, TILE_M, TILE_N, TILE_K, HAS_BIAS, SWIZZLE) + + +@ct.kernel(occupancy=4) +def _row_max_kernel( + input, + output, + n_cols, + BLOCK_SIZE: ConstInt, +): + row = ct.bid(0) + num_chunks = (n_cols + BLOCK_SIZE - 1) // BLOCK_SIZE + row_max_tile = ct.full((1,), -math.inf, dtype=ct.float32) + + for chunk in range(num_chunks): + columns = ct.arange(BLOCK_SIZE, dtype=ct.int32) + chunk * BLOCK_SIZE + values = ct.astype( + ct.gather( + input, + (row, columns), + check_bounds=True, + padding_value=-math.inf, + latency=3, + ), + ct.float32, + ) + row_max = ct.maximum( + ct.sum(row_max_tile, 0, keepdims=False), + ct.max(values, 0, keepdims=False), + ) + row_max_tile = ct.full((1,), row_max, dtype=ct.float32) + + ct.scatter(output, row, ct.sum(row_max_tile, 0, keepdims=False)) + + +@ct.kernel(occupancy=4) +def _vocab_parallel_ce_forward_kernel( + logits, + logits_max, + target, + predicted_logit, + sum_exp, + vocab_start, + n_cols, + ignore_index, + BLOCK_SIZE: ConstInt, +): + row = ct.bid(0) + y_global = ct.load(target, row, shape=()) + maximum = ct.astype(ct.load(logits_max, row, shape=()), ct.float32) + is_ignored = y_global == ignore_index + target_off_rank = (y_global < vocab_start) or (y_global >= vocab_start + n_cols) + y_local = ct.astype(y_global - vocab_start, ct.int32) + + if is_ignored or target_off_rank: + predicted = 0.0 + else: + target_index = ct.add(ct.arange(1, dtype=ct.int32), y_local) + target_tile = ct.gather(logits, (row, target_index), check_bounds=False) + predicted = ct.sum(ct.astype(target_tile, ct.float32), 0, keepdims=False) - maximum + + num_chunks = (n_cols + BLOCK_SIZE - 1) // BLOCK_SIZE + sum_exp_tile = ct.full((1,), 0.0, dtype=ct.float32) + for chunk in range(num_chunks): + columns = ct.arange(BLOCK_SIZE, dtype=ct.int32) + chunk * BLOCK_SIZE + in_bounds = columns < n_cols + values = ct.astype( + ct.gather( + logits, + (row, columns), + check_bounds=True, + padding_value=-math.inf, + latency=3, + ), + ct.float32, + ) + exponentials = ct.exp2((values - maximum) * LOG2E, flush_to_zero=True) + exponentials = ct.where(in_bounds, exponentials, 0.0) + running_sum = ct.sum(sum_exp_tile, 0, keepdims=False) + sum_exp_tile = ct.full( + (1,), + running_sum + ct.sum(exponentials, 0, keepdims=False), + dtype=ct.float32, + ) + + ct.scatter(predicted_logit, row, predicted) + ct.scatter(sum_exp, row, ct.sum(sum_exp_tile, 0, keepdims=False)) + + +@ct.kernel(occupancy=4) +def _vocab_parallel_ce_backward_kernel( + logits, + logits_max, + sum_exp, + target, + grad_output, + vocab_start, + n_cols, + ignore_index, + BLOCK_SIZE: ConstInt, +): + row = ct.bid(0) + y_global = ct.load(target, row, shape=()) + num_chunks = (n_cols + BLOCK_SIZE - 1) // BLOCK_SIZE + + if y_global == ignore_index: + for chunk in range(num_chunks): + columns = ct.arange(BLOCK_SIZE, dtype=ct.int32) + chunk * BLOCK_SIZE + zeros = ct.full((BLOCK_SIZE,), 0.0, dtype=logits.dtype) + ct.scatter(logits, (row, columns), zeros, check_bounds=True) + return + + target_off_rank = (y_global < vocab_start) or (y_global >= vocab_start + n_cols) + y_local = ct.astype(y_global - vocab_start, ct.int32) + maximum = ct.astype(ct.load(logits_max, row, shape=()), ct.float32) + global_sum = ct.astype(ct.load(sum_exp, row, shape=()), ct.float32) + upstream = ct.astype(ct.load(grad_output, row, shape=()), ct.float32) + + for chunk in range(num_chunks): + columns = ct.arange(BLOCK_SIZE, dtype=ct.int32) + chunk * BLOCK_SIZE + values = ct.astype( + ct.gather(logits, (row, columns), check_bounds=True, padding_value=-math.inf), + ct.float32, + ) + exponentials = ct.exp2((values - maximum) * LOG2E, flush_to_zero=True) + gradient = exponentials / global_sum + if not target_off_rank: + gradient = ct.where(columns == y_local, gradient - 1.0, gradient) + gradient = gradient * upstream + ct.scatter( + logits, + (row, columns), + ct.astype(gradient, logits.dtype), + check_bounds=True, + ) + + +@ct.kernel(occupancy=4) +def _loss_kernel( + sum_exp, + predicted_logit, + target, + output, + ignore_index, +): + row = ct.bid(0) + y_global = ct.load(target, row, shape=()) + if y_global == ignore_index: + loss = 0.0 + else: + denominator = ct.astype(ct.load(sum_exp, row, shape=()), ct.float32) + predicted = ct.astype(ct.load(predicted_logit, row, shape=()), ct.float32) + loss = ct.log(denominator) - predicted + ct.scatter(output, row, loss) + + +@ct.kernel(occupancy=4) +def _column_sum_kernel( + input, + output, + n_rows, + BLOCK_SIZE: ConstInt, +): + column = ct.bid(0) + num_chunks = (n_rows + BLOCK_SIZE - 1) // BLOCK_SIZE + total_tile = ct.full((1,), 0.0, dtype=ct.float32) + + for chunk in range(num_chunks): + rows = ct.arange(BLOCK_SIZE, dtype=ct.int32) + chunk * BLOCK_SIZE + values = ct.astype( + ct.gather(input, (rows, column), check_bounds=True, padding_value=0.0), + ct.float32, + ) + running_total = ct.sum(total_tile, 0, keepdims=False) + total_tile = ct.full( + (1,), + running_total + ct.sum(values, 0, keepdims=False), + dtype=ct.float32, + ) + + ct.scatter(output, column, ct.astype(ct.sum(total_tile, 0, keepdims=False), output.dtype)) + + +def _select_row_block_size(size: int) -> int: + return min(MAX_ROW_BLOCK_SIZE, _next_power_of_2(size)) + + +def _cutile_matmul( + a: torch.Tensor, + b: torch.Tensor, + *, + operation: str, + bias: torch.Tensor | None = None, + output_dtype: torch.dtype | None = None, +) -> torch.Tensor: + if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]: + raise ValueError(f"matmul expects [M, K] @ [K, N], got {tuple(a.shape)} and {tuple(b.shape)}.") + + output = torch.empty( + (a.shape[0], b.shape[1]), + device=a.device, + dtype=output_dtype or a.dtype, + ) + if operation == "projection": + kernel, tile = ( + (_matmul_2cta_kernel, (256, 256, 128)) if a.shape[0] <= 1024 else (_matmul_2cta_kernel, (512, 256, 64)) + ) + elif operation == "dx": + if a.shape[0] <= 1024: + kernel, tile = _matmul_1cta_kernel, (128, 128, 64) + elif a.shape[1] > 16000 or a.shape[0] > 16384: + kernel, tile = _matmul_2cta_kernel, (512, 256, 64) + else: + kernel, tile = _matmul_1cta_kernel, (256, 256, 64) + elif operation == "dw": + if a.shape[1] >= 16384 and (a.shape[0] > 16000 or a.shape[1] == 16384): + kernel, tile = _matmul_2cta_kernel, (512, 256, 64) + else: + use_single_cta = a.shape[1] <= 1024 or a.shape[0] > 16000 or a.shape[1] > 16384 + kernel = _matmul_1cta_kernel if use_single_cta else _matmul_2cta_kernel + tile = (256, 256, 64) + else: + raise ValueError(f"unknown FLCE GEMM operation: {operation!r}.") + + tile_m, tile_n, tile_k = tile + swizzle = operation == "projection" and b.shape[1] > 16000 + grid = ( + ct.cdiv(a.shape[0], tile_m) * ct.cdiv(b.shape[1], tile_n), + 1, + 1, + ) + ct.launch( + torch.cuda.current_stream(), + grid, + kernel, + ( + a, + b, + bias if bias is not None else output, + output, + tile_m, + tile_n, + tile_k, + bias is not None, + swizzle, + ), + ) + return output + + +def _cutile_row_max(input: torch.Tensor) -> torch.Tensor: + output = torch.empty(input.shape[0], device=input.device, dtype=torch.float32) + block_size = 16384 if input.shape[1] > 16384 else _select_row_block_size(input.shape[1]) + ct.launch( + torch.cuda.current_stream(), + (input.shape[0], 1, 1), + _row_max_kernel, + (input, output, int(input.shape[1]), int(block_size)), + ) + return output + + +def _cutile_ce_forward( + logits: torch.Tensor, + logits_max: torch.Tensor, + target: torch.Tensor, + vocab_start: int, + ignore_index: int, +) -> torch.Tensor: + rows, vocab_local = logits.shape + stats = torch.empty((2, rows), device=logits.device, dtype=torch.float32) + predicted_logit = stats[0] + sum_exp = stats[1] + block_size = 16384 if vocab_local > 16384 else _select_row_block_size(vocab_local) + ct.launch( + torch.cuda.current_stream(), + (rows, 1, 1), + _vocab_parallel_ce_forward_kernel, + ( + logits, + logits_max, + target, + predicted_logit, + sum_exp, + int(vocab_start), + int(vocab_local), + int(ignore_index), + int(block_size), + ), + ) + return stats + + +def _cutile_ce_backward( + logits: torch.Tensor, + logits_max: torch.Tensor, + sum_exp: torch.Tensor, + target: torch.Tensor, + grad_output: torch.Tensor, + vocab_start: int, + ignore_index: int, +) -> None: + block_size = min(2048, _select_row_block_size(logits.shape[1])) + ct.launch( + torch.cuda.current_stream(), + (logits.shape[0], 1, 1), + _vocab_parallel_ce_backward_kernel, + ( + logits, + logits_max, + sum_exp, + target, + grad_output, + int(vocab_start), + int(logits.shape[1]), + int(ignore_index), + int(block_size), + ), + ) + + +def _cutile_loss( + sum_exp: torch.Tensor, + predicted_logit: torch.Tensor, + target: torch.Tensor, + ignore_index: int, +) -> torch.Tensor: + output = torch.empty_like(sum_exp) + ct.launch( + torch.cuda.current_stream(), + (target.numel(), 1, 1), + _loss_kernel, + (sum_exp, predicted_logit, target, output, int(ignore_index)), + ) + return output + + +def _cutile_column_sum(input: torch.Tensor, output_dtype: torch.dtype) -> torch.Tensor: + output = torch.empty(input.shape[1], device=input.device, dtype=output_dtype) + block_size = _select_row_block_size(input.shape[0]) + ct.launch( + torch.cuda.current_stream(), + (input.shape[1], 1, 1), + _column_sum_kernel, + (input, output, int(input.shape[0]), int(block_size)), + ) + return output + + +def _materialized_backward(ctx, grad_output: torch.Tensor): + hidden, weight, logits, logits_max, sum_exp, target = ctx.saved_tensors + grad_output_1d = grad_output.contiguous().reshape(-1).float() + _cutile_ce_backward( + logits, + logits_max, + sum_exp, + target, + grad_output_1d, + ctx.vocab_start, + ctx.ignore_index, + ) + + grad_hidden = _cutile_matmul( + logits, + weight, + operation="dx", + output_dtype=torch.float32 if logits.shape[0] <= 1024 else None, + ).to(ctx.hidden_dtype) + reduce_work = ( + dist.all_reduce( + grad_hidden, + op=dist.ReduceOp.SUM, + group=ctx.tp_group, + async_op=True, + ) + if ctx.tp_world > 1 + else None + ) + grad_weight = _cutile_matmul(logits.t(), hidden, operation="dw") + grad_bias = _cutile_column_sum(logits, ctx.bias_dtype) if ctx.has_bias else None + + if reduce_work is not None: + reduce_work.wait() + grad_hidden = grad_hidden.reshape(ctx.original_hidden_shape) + return grad_hidden, grad_weight, grad_bias + + +class LigerMegatronFusedLinearCrossEntropyFunction(torch.autograd.Function): + """Hidden-to-loss tensor-parallel FLCE using CuTile local kernels.""" + + @staticmethod + def forward( + ctx, + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None, + tp_group, + ignore_index: int, + ) -> torch.Tensor: + _validate_megatron_flce_inputs(hidden, weight, target, bias) + if hidden.device.type != "cuda" or hidden.dtype not in (torch.bfloat16, torch.float16): + raise RuntimeError("CuTile Megatron FLCE requires a CUDA GPU and float16 or bfloat16 inputs.") + + tp_rank, tp_world = _tp_rank_and_world(tp_group) + vocab_local = weight.shape[0] + vocab_global = vocab_local * tp_world + vocab_start = tp_rank * vocab_local + + flat_target = target.reshape(-1).contiguous() + valid = flat_target != ignore_index + invalid = valid & ((flat_target < 0) | (flat_target >= vocab_global)) + valid_targets = ~torch.any(invalid) + if hasattr(torch, "_assert_async"): + torch._assert_async(valid_targets, f"non-ignored targets must be in [0, {vocab_global}).") + elif not valid_targets.item(): + raise ValueError(f"non-ignored targets must be in [0, {vocab_global}).") + + original_hidden_shape = hidden.shape + hidden_2d = hidden.reshape(-1, hidden.shape[-1]).contiguous() + weight_2d = weight.contiguous() + bias_1d = bias.contiguous() if bias is not None else None + + logits = _cutile_matmul(hidden_2d, weight_2d.t(), operation="projection", bias=bias_1d) + logits_max = _cutile_row_max(logits) + if tp_world > 1: + dist.all_reduce(logits_max, op=dist.ReduceOp.MAX, group=tp_group) + + stats = _cutile_ce_forward( + logits, + logits_max, + flat_target, + vocab_start, + ignore_index, + ) + if tp_world > 1: + dist.all_reduce(stats, op=dist.ReduceOp.SUM, group=tp_group) + predicted_logit = stats[0] + sum_exp = stats[1] + + loss = _cutile_loss(sum_exp, predicted_logit, flat_target, ignore_index) + + ctx.save_for_backward(hidden_2d, weight_2d, logits, logits_max, sum_exp, flat_target) + ctx.has_bias = bias is not None + ctx.bias_dtype = bias.dtype if bias is not None else None + ctx.tp_group = tp_group + ctx.tp_world = tp_world + ctx.vocab_start = vocab_start + ctx.ignore_index = ignore_index + ctx.original_hidden_shape = original_hidden_shape + ctx.hidden_dtype = hidden.dtype + return loss.reshape(target.shape) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + grad_hidden, grad_weight, grad_bias = _materialized_backward(ctx, grad_output) + return grad_hidden, grad_weight, None, grad_bias, None, None + + +def liger_megatron_fused_linear_cross_entropy( + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None = None, + tp_group=None, + ignore_index: int = -100, +) -> torch.Tensor: + """Compute Megatron FLCE with CuTile local kernels and NCCL TP collectives.""" + return LigerMegatronFusedLinearCrossEntropyFunction.apply( + hidden, + weight, + target, + bias, + tp_group, + ignore_index, + ) diff --git a/src/liger_kernel/ops/megatron_fused_linear_cross_entropy.py b/src/liger_kernel/ops/megatron_fused_linear_cross_entropy.py new file mode 100644 index 000000000..e4c9ed9d6 --- /dev/null +++ b/src/liger_kernel/ops/megatron_fused_linear_cross_entropy.py @@ -0,0 +1,725 @@ +"""Portable all-Triton tensor-parallel fused linear cross entropy for Megatron. + +Each tensor-parallel rank owns a contiguous vocabulary shard. Forward performs +one Triton projection GEMM, computes globally normalized cross entropy, and +saves the local logits in the projection dtype. Backward converts that buffer +to dlogits in-place before Triton dX and dW GEMMs. Tensor-parallel collectives +remain NCCL/RCCL calls between architecture-independent kernels. +""" + +from __future__ import annotations + +import torch +import torch.distributed as dist +import triton +import triton.language as tl + + +def _matmul_autotune_configs(): + return [ + triton.Config( + {"BLOCK_SIZE_M": 256, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 64, "GROUP_SIZE_M": 8}, + num_stages=3, + num_warps=8, + ), + triton.Config( + {"BLOCK_SIZE_M": 256, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 64, "GROUP_SIZE_M": 8}, + num_stages=3, + num_warps=8, + ), + triton.Config( + {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 128, "GROUP_SIZE_M": 8}, + num_stages=3, + num_warps=8, + ), + triton.Config( + {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 64, "GROUP_SIZE_M": 8}, + num_stages=3, + num_warps=8, + ), + triton.Config( + {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 256, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, + num_stages=4, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, + num_stages=4, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 128, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, + num_stages=4, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 128, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, + num_stages=4, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 64, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 8}, + num_stages=3, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 64, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 4}, + num_stages=3, + num_warps=4, + ), + triton.Config( + {"BLOCK_SIZE_M": 32, "BLOCK_SIZE_N": 32, "BLOCK_SIZE_K": 32, "GROUP_SIZE_M": 4}, + num_stages=3, + num_warps=4, + ), + ] + + +def _split_k_matmul_autotune_configs(): + return [ + triton.Config( + { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 2, + }, + num_stages=3, + num_warps=4, + ), + triton.Config( + { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 4, + }, + num_stages=3, + num_warps=4, + ), + triton.Config( + { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 8, + }, + num_stages=3, + num_warps=4, + ), + triton.Config( + { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 4, + }, + num_stages=3, + num_warps=8, + ), + triton.Config( + { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 4, + }, + num_stages=3, + num_warps=8, + ), + triton.Config( + { + "BLOCK_SIZE_M": 128, + "BLOCK_SIZE_N": 256, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 8, + "SPLIT_K": 4, + }, + num_stages=3, + num_warps=8, + ), + ] + + +@triton.autotune(configs=_matmul_autotune_configs(), key=["M", "N", "K"]) +@triton.jit +def _matmul_kernel( + a_ptr, + b_ptr, + bias_ptr, + output_ptr, + M, + N, + K, + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_om, + stride_on, + HAS_BIAS: tl.constexpr, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, +): + pid = tl.program_id(0) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m + pid_n = (pid % num_pid_in_group) // group_size_m + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k_start in range(0, tl.cdiv(K, BLOCK_SIZE_K)): + k_remaining = K - k_start * BLOCK_SIZE_K + a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < k_remaining), other=0.0) + b = tl.load(b_ptrs, mask=(offs_k[:, None] < k_remaining) & (offs_n[None, :] < N), other=0.0) + accumulator = tl.dot(a, b, accumulator) + a_ptrs += BLOCK_SIZE_K * stride_ak + b_ptrs += BLOCK_SIZE_K * stride_bk + + if HAS_BIAS: + bias = tl.load(bias_ptr + offs_n, mask=offs_n < N, other=0.0).to(tl.float32) + accumulator += bias[None, :] + + output_ptrs = output_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on + tl.store(output_ptrs, accumulator, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + + +@triton.autotune( + configs=_split_k_matmul_autotune_configs(), + key=["M", "N", "K"], + reset_to_zero=["output_ptr"], +) +@triton.jit +def _split_k_matmul_kernel( + a_ptr, + b_ptr, + output_ptr, + M, + N, + K, + stride_am, + stride_ak, + stride_bk, + stride_bn, + stride_om, + stride_on, + BLOCK_SIZE_M: tl.constexpr, + BLOCK_SIZE_N: tl.constexpr, + BLOCK_SIZE_K: tl.constexpr, + GROUP_SIZE_M: tl.constexpr, + SPLIT_K: tl.constexpr, +): + pid = tl.program_id(0) + split_k_id = tl.program_id(1) + num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) + num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) + num_pid_in_group = GROUP_SIZE_M * num_pid_n + group_id = pid // num_pid_in_group + first_pid_m = group_id * GROUP_SIZE_M + group_size_m = min(num_pid_m - first_pid_m, GROUP_SIZE_M) + pid_m = first_pid_m + (pid % num_pid_in_group) % group_size_m + pid_n = (pid % num_pid_in_group) // group_size_m + + offs_m = pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M) + offs_n = pid_n * BLOCK_SIZE_N + tl.arange(0, BLOCK_SIZE_N) + offs_k = split_k_id * BLOCK_SIZE_K + tl.arange(0, BLOCK_SIZE_K) + a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak + b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn + + accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) + for k_start in range(0, tl.cdiv(K, BLOCK_SIZE_K * SPLIT_K)): + k_remaining = K - k_start * BLOCK_SIZE_K * SPLIT_K + a = tl.load(a_ptrs, mask=(offs_m[:, None] < M) & (offs_k[None, :] < k_remaining), other=0.0) + b = tl.load(b_ptrs, mask=(offs_k[:, None] < k_remaining) & (offs_n[None, :] < N), other=0.0) + accumulator = tl.dot(a, b, accumulator) + a_ptrs += BLOCK_SIZE_K * SPLIT_K * stride_ak + b_ptrs += BLOCK_SIZE_K * SPLIT_K * stride_bk + + output_ptrs = output_ptr + offs_m[:, None] * stride_om + offs_n[None, :] * stride_on + tl.atomic_add(output_ptrs, accumulator, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) + + +@triton.jit +def _row_max_kernel( + input_ptr, + output_ptr, + n_cols, + input_row_stride, + BLOCK_SIZE: tl.constexpr, +): + row = tl.program_id(0).to(tl.int64) + row_ptr = input_ptr + row * input_row_stride + row_max = -float("inf") + for start in range(0, n_cols, BLOCK_SIZE): + offsets = start + tl.arange(0, BLOCK_SIZE) + values = tl.load(row_ptr + offsets, mask=offsets < n_cols, other=-float("inf")).to(tl.float32) + row_max = tl.maximum(row_max, tl.max(values)) + tl.store(output_ptr + row, row_max) + + +@triton.jit +def _ce_forward_stats_kernel( + logits_ptr, + logits_stride, + logits_max_ptr, + target_ptr, + predicted_logit_ptr, + sum_exp_ptr, + vocab_start, + n_cols, + ignore_index, + BLOCK_SIZE: tl.constexpr, +): + row = tl.program_id(0).to(tl.int64) + row_ptr = logits_ptr + row * logits_stride + target = tl.load(target_ptr + row) + maximum = tl.load(logits_max_ptr + row).to(tl.float32) + target_off_rank = (target < vocab_start) | (target >= vocab_start + n_cols) + + if target == ignore_index or target_off_rank: + predicted_logit = 0.0 + else: + predicted_logit = tl.load(row_ptr + target - vocab_start).to(tl.float32) - maximum + + sum_exp = 0.0 + for start in range(0, n_cols, BLOCK_SIZE): + offsets = start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_cols + logits = tl.load(row_ptr + offsets, mask=mask, other=-float("inf")).to(tl.float32) + sum_exp += tl.sum(tl.exp(logits - maximum)) + + tl.store(predicted_logit_ptr + row, predicted_logit) + tl.store(sum_exp_ptr + row, sum_exp) + + +@triton.jit +def _ce_backward_from_logits_kernel( + logits_ptr, + logits_stride, + logits_max_ptr, + sum_exp_ptr, + target_ptr, + grad_output_ptr, + vocab_start, + n_cols, + ignore_index, + BLOCK_SIZE: tl.constexpr, +): + row = tl.program_id(0).to(tl.int64) + row_ptr = logits_ptr + row * logits_stride + target = tl.load(target_ptr + row) + + if target == ignore_index: + for start in range(0, n_cols, BLOCK_SIZE): + offsets = start + tl.arange(0, BLOCK_SIZE) + tl.store(row_ptr + offsets, 0.0, mask=offsets < n_cols) + return + + maximum = tl.load(logits_max_ptr + row).to(tl.float32) + sum_exp = tl.load(sum_exp_ptr + row).to(tl.float32) + grad_output = tl.load(grad_output_ptr + row).to(tl.float32) + target_off_rank = (target < vocab_start) | (target >= vocab_start + n_cols) + target_local = target - vocab_start + + for start in range(0, n_cols, BLOCK_SIZE): + offsets = start + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_cols + logits = tl.load(row_ptr + offsets, mask=mask, other=-float("inf")).to(tl.float32) + gradient = tl.exp(logits - maximum) / sum_exp + if not target_off_rank: + gradient = tl.where(offsets == target_local, gradient - 1.0, gradient) + tl.store(row_ptr + offsets, gradient * grad_output, mask=mask) + + +@triton.jit +def _loss_kernel( + sum_exp_ptr, + predicted_logit_ptr, + target_ptr, + loss_ptr, + n_rows, + ignore_index, + BLOCK_SIZE: tl.constexpr, +): + offsets = tl.program_id(0).to(tl.int64) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < n_rows + sum_exp = tl.load(sum_exp_ptr + offsets, mask=mask, other=1.0) + predicted_logit = tl.load(predicted_logit_ptr + offsets, mask=mask, other=0.0) + target = tl.load(target_ptr + offsets, mask=mask, other=ignore_index) + loss = tl.log(sum_exp) - predicted_logit + loss = tl.where(target == ignore_index, 0.0, loss) + tl.store(loss_ptr + offsets, loss, mask=mask) + + +@triton.jit +def _column_sum_kernel( + input_ptr, + output_ptr, + n_rows, + input_row_stride, + BLOCK_SIZE: tl.constexpr, +): + col = tl.program_id(0).to(tl.int64) + offsets = tl.arange(0, BLOCK_SIZE) + total = 0.0 + for start in range(0, n_rows, BLOCK_SIZE): + rows = start + offsets + values = tl.load(input_ptr + rows * input_row_stride + col, mask=rows < n_rows, other=0.0) + total += tl.sum(values.to(tl.float32)) + tl.store(output_ptr + col, total) + + +def _triton_matmul( + a: torch.Tensor, + b: torch.Tensor, + *, + bias: torch.Tensor | None = None, + output_dtype: torch.dtype | None = None, +) -> torch.Tensor: + if a.ndim != 2 or b.ndim != 2 or a.shape[1] != b.shape[0]: + raise ValueError(f"matmul expects [M, K] @ [K, N], got {tuple(a.shape)} and {tuple(b.shape)}.") + m, k = a.shape + n = b.shape[1] + output = torch.empty((m, n), device=a.device, dtype=output_dtype or a.dtype) + grid = lambda meta: (triton.cdiv(m, meta["BLOCK_SIZE_M"]) * triton.cdiv(n, meta["BLOCK_SIZE_N"]),) + _matmul_kernel[grid]( + a, + b, + bias if bias is not None else output, + output, + m, + n, + k, + a.stride(0), + a.stride(1), + b.stride(0), + b.stride(1), + output.stride(0), + output.stride(1), + HAS_BIAS=bias is not None, + ) + return output + + +def _triton_dx_matmul(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + m, k = a.shape + n = b.shape[1] + if m > 1024 or k < 4096: + return _triton_matmul(a, b) + + output = torch.zeros((m, n), device=a.device, dtype=torch.float32) + grid = lambda meta: ( + triton.cdiv(m, meta["BLOCK_SIZE_M"]) * triton.cdiv(n, meta["BLOCK_SIZE_N"]), + meta["SPLIT_K"], + ) + _split_k_matmul_kernel[grid]( + a, + b, + output, + m, + n, + k, + a.stride(0), + a.stride(1), + b.stride(0), + b.stride(1), + output.stride(0), + output.stride(1), + ) + return output + + +def _triton_row_max(input: torch.Tensor, block_size: int) -> torch.Tensor: + from liger_kernel.ops.vocab_parallel_cross_entropy import _get_num_warps + + output = torch.empty(input.shape[0], device=input.device, dtype=torch.float32) + _row_max_kernel[(input.shape[0],)]( + input, + output, + input.shape[1], + input.stride(0), + BLOCK_SIZE=block_size, + num_warps=_get_num_warps(block_size), + ) + return output + + +def _ce_forward_stats( + logits: torch.Tensor, + logits_max: torch.Tensor, + target: torch.Tensor, + vocab_start: int, + ignore_index: int, + block_size: int, +) -> torch.Tensor: + from liger_kernel.ops.vocab_parallel_cross_entropy import _get_num_warps + + stats = torch.empty((2, logits.shape[0]), device=logits.device, dtype=torch.float32) + _ce_forward_stats_kernel[(logits.shape[0],)]( + logits, + logits.stride(0), + logits_max, + target, + stats[0], + stats[1], + vocab_start, + logits.shape[1], + ignore_index, + BLOCK_SIZE=block_size, + num_warps=_get_num_warps(block_size), + ) + return stats + + +def _ce_backward_from_logits( + logits: torch.Tensor, + logits_max: torch.Tensor, + sum_exp: torch.Tensor, + target: torch.Tensor, + grad_output: torch.Tensor, + vocab_start: int, + ignore_index: int, + block_size: int, +) -> None: + from liger_kernel.ops.vocab_parallel_cross_entropy import _get_num_warps + + _ce_backward_from_logits_kernel[(logits.shape[0],)]( + logits, + logits.stride(0), + logits_max, + sum_exp, + target, + grad_output, + vocab_start, + logits.shape[1], + ignore_index, + BLOCK_SIZE=block_size, + num_warps=_get_num_warps(block_size), + ) + + +def _triton_loss( + sum_exp: torch.Tensor, + predicted_logit: torch.Tensor, + target: torch.Tensor, + ignore_index: int, +) -> torch.Tensor: + output = torch.empty_like(sum_exp) + block_size = 256 + _loss_kernel[(triton.cdiv(target.numel(), block_size),)]( + sum_exp, + predicted_logit, + target, + output, + target.numel(), + ignore_index, + BLOCK_SIZE=block_size, + num_warps=4, + ) + return output + + +def _triton_column_sum(input: torch.Tensor) -> torch.Tensor: + from liger_kernel.ops.vocab_parallel_cross_entropy import _get_num_warps + from liger_kernel.ops.vocab_parallel_cross_entropy import _select_block_size + + block_size = _select_block_size(input.shape[0]) + output = torch.empty(input.shape[1], device=input.device, dtype=torch.float32) + _column_sum_kernel[(input.shape[1],)]( + input, + output, + input.shape[0], + input.stride(0), + BLOCK_SIZE=block_size, + num_warps=_get_num_warps(block_size), + ) + return output + + +def _tp_rank_and_world(tp_group) -> tuple[int, int]: + if tp_group is None: + return 0, 1 + world = dist.get_world_size(tp_group) + if world == 1: + return 0, 1 + return dist.get_rank(tp_group), world + + +def _validate_megatron_flce_inputs( + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None, +) -> None: + if hidden.ndim < 2: + raise ValueError(f"hidden must have at least 2 dimensions, got shape {tuple(hidden.shape)}.") + if weight.ndim != 2: + raise ValueError(f"weight must be 2-D [V_local, H], got shape {tuple(weight.shape)}.") + if tuple(target.shape) != tuple(hidden.shape[:-1]): + raise ValueError( + f"target shape must equal hidden.shape[:-1]; got target={tuple(target.shape)}, " + f"hidden={tuple(hidden.shape)}." + ) + if target.dtype != torch.long: + raise TypeError(f"target must have dtype torch.long, got {target.dtype}.") + if hidden.shape[-1] != weight.shape[1]: + raise ValueError(f"hidden size mismatch: hidden has H={hidden.shape[-1]}, weight has H={weight.shape[1]}.") + if target.numel() == 0 or hidden.shape[-1] == 0 or weight.shape[0] == 0: + raise ValueError("hidden, weight, and target dimensions must be non-empty.") + if hidden.dtype != weight.dtype: + raise TypeError(f"hidden and weight must have the same dtype, got {hidden.dtype} and {weight.dtype}.") + if hidden.device != weight.device or hidden.device != target.device: + raise ValueError("hidden, weight, and target must be on the same device.") + if bias is not None: + if bias.ndim != 1 or bias.shape[0] != weight.shape[0]: + raise ValueError(f"bias must have shape ({weight.shape[0]},), got {tuple(bias.shape)}.") + if bias.device != hidden.device or bias.dtype != hidden.dtype: + raise TypeError("bias must have the same device and dtype as hidden.") + + +def _materialized_backward(ctx, grad_output: torch.Tensor): + """Convert saved logits to dlogits and form projection gradients.""" + hidden, weight, logits, logits_max, sum_exp_global, target = ctx.saved_tensors + grad_out = grad_output.contiguous().reshape(-1).float() + _ce_backward_from_logits( + logits, + logits_max, + sum_exp_global, + target, + grad_out, + ctx.vocab_start, + ctx.ignore_index, + ctx.ce_block_size, + ) + + grad_hidden = _triton_dx_matmul(logits, weight).to(ctx.hidden_dtype) + reduce_work = ( + dist.all_reduce( + grad_hidden, + op=dist.ReduceOp.SUM, + group=ctx.tp_group, + async_op=True, + ) + if ctx.tp_world > 1 + else None + ) + grad_weight = _triton_matmul(logits.t(), hidden) + grad_bias = _triton_column_sum(logits).to(ctx.bias_dtype) if ctx.has_bias else None + + if reduce_work is not None: + reduce_work.wait() + grad_hidden = grad_hidden.reshape(ctx.original_hidden_shape) + return grad_hidden, grad_weight, grad_bias + + +class LigerMegatronFusedLinearCrossEntropyFunction(torch.autograd.Function): + """Hidden-to-loss tensor-parallel FLCE with saved low-precision CE state.""" + + @staticmethod + def forward( + ctx, + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None, + tp_group, + ignore_index: int, + ) -> torch.Tensor: + _validate_megatron_flce_inputs(hidden, weight, target, bias) + if hidden.device.type != "cuda" or hidden.dtype not in (torch.bfloat16, torch.float16): + raise RuntimeError("Megatron FLCE requires a CUDA GPU and float16 or bfloat16 inputs.") + + tp_rank, tp_world = _tp_rank_and_world(tp_group) + vocab_local = weight.shape[0] + vocab_global = vocab_local * tp_world + vocab_start = tp_rank * vocab_local + + flat_target = target.reshape(-1).contiguous() + valid = flat_target != ignore_index + invalid = valid & ((flat_target < 0) | (flat_target >= vocab_global)) + valid_targets = ~torch.any(invalid) + if hasattr(torch, "_assert_async"): + torch._assert_async(valid_targets, f"non-ignored targets must be in [0, {vocab_global}).") + elif not valid_targets.item(): + raise ValueError(f"non-ignored targets must be in [0, {vocab_global}).") + + original_hidden_shape = hidden.shape + hidden_2d = hidden.reshape(-1, hidden.shape[-1]).contiguous() + weight_2d = weight.contiguous() + bias_1d = bias.contiguous() if bias is not None else None + + from liger_kernel.ops.vocab_parallel_cross_entropy import _select_block_size + + logits = _triton_matmul(hidden_2d, weight_2d.t(), bias=bias_1d) + ce_block_size = _select_block_size(vocab_local) + logits_max = _triton_row_max(logits, ce_block_size) + if tp_world > 1: + dist.all_reduce(logits_max, op=dist.ReduceOp.MAX, group=tp_group) + + stats = _ce_forward_stats( + logits, + logits_max, + flat_target, + vocab_start, + ignore_index, + ce_block_size, + ) + predicted_logit = stats[0] + sum_exp = stats[1] + if tp_world > 1: + dist.all_reduce(stats, op=dist.ReduceOp.SUM, group=tp_group) + + loss = _triton_loss(sum_exp, predicted_logit, flat_target, ignore_index) + + ctx.save_for_backward(hidden_2d, weight_2d, logits, logits_max, sum_exp, flat_target) + ctx.has_bias = bias is not None + ctx.bias_dtype = bias.dtype if bias is not None else None + ctx.tp_group = tp_group + ctx.tp_world = tp_world + ctx.vocab_start = vocab_start + ctx.ignore_index = ignore_index + ctx.ce_block_size = ce_block_size + ctx.original_hidden_shape = original_hidden_shape + ctx.hidden_dtype = hidden.dtype + return loss.reshape(target.shape) + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): + grad_hidden, grad_weight, grad_bias = _materialized_backward(ctx, grad_output) + return grad_hidden, grad_weight, None, grad_bias, None, None + + +def liger_megatron_fused_linear_cross_entropy( + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None = None, + tp_group=None, + ignore_index: int = -100, +) -> torch.Tensor: + """Compute per-token loss from replicated hidden states and a local vocab shard.""" + return LigerMegatronFusedLinearCrossEntropyFunction.apply( + hidden, + weight, + target, + bias, + tp_group, + ignore_index, + ) diff --git a/test/megatron/test_cutedsl_fused_linear_cross_entropy.py b/test/megatron/test_cutedsl_fused_linear_cross_entropy.py new file mode 100644 index 000000000..40ff34f5a --- /dev/null +++ b/test/megatron/test_cutedsl_fused_linear_cross_entropy.py @@ -0,0 +1,155 @@ +from __future__ import annotations + +import os +import subprocess +import sys +import tempfile + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch.nn.functional as F + +pytest.importorskip("cutlass.cute") + +from liger_kernel.ops.cutedsl.ops.megatron_fused_linear_cross_entropy import ( # noqa: E402 + liger_megatron_fused_linear_cross_entropy, +) + +pytestmark = [ + pytest.mark.skipif(not torch.cuda.is_available(), reason="CuTe DSL FLCE requires CUDA"), + pytest.mark.skipif( + torch.cuda.is_available() and torch.cuda.get_device_capability()[0] < 10, + reason="native CuTe DSL Megatron FLCE requires Blackwell", + ), +] + + +def _reference_loss(hidden, weight, target, bias=None, ignore_index=-100): + logits = hidden.float() @ weight.float().t() + if bias is not None: + logits = logits + bias.float() + return F.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + target.reshape(-1), + reduction="none", + ignore_index=ignore_index, + ).reshape(target.shape) + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("with_bias", [False, True]) +def test_cutedsl_megatron_flce_tp1_matches_pytorch(dtype, with_bias): + torch.manual_seed(42) + hidden_base = torch.randn(3, 2, 65, device="cuda", dtype=dtype) + weight_base = torch.randn(33, 65, device="cuda", dtype=dtype) * 0.02 + bias_base = torch.randn(33, device="cuda", dtype=dtype) * 0.02 if with_bias else None + target = torch.randint(0, 33, (3, 2), device="cuda") + target[0, 0] = -100 + upstream = torch.randn(3, 2, device="cuda") + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_base.clone().requires_grad_(True) + bias_ref = bias_base.clone().requires_grad_(True) if bias_base is not None else None + hidden_cutedsl = hidden_base.clone().requires_grad_(True) + weight_cutedsl = weight_base.clone().requires_grad_(True) + bias_cutedsl = bias_base.clone().requires_grad_(True) if bias_base is not None else None + + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + actual = liger_megatron_fused_linear_cross_entropy( + hidden_cutedsl, + weight_cutedsl, + target, + bias=bias_cutedsl, + ) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + reference.backward(upstream) + actual.backward(upstream) + torch.testing.assert_close(hidden_cutedsl.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_cutedsl.grad, weight_ref.grad, atol=5e-3, rtol=5e-2) + if with_bias: + torch.testing.assert_close(bias_cutedsl.grad, bias_ref.grad, atol=5e-3, rtol=5e-2) + + +def test_cutedsl_megatron_flce_backend_dispatch(): + env = os.environ.copy() + env["LIGER_KERNEL_IMPL"] = "cutedsl" + result = subprocess.run( + [ + sys.executable, + "-c", + ( + "from liger_kernel.megatron.fused_linear_cross_entropy import " + "LigerMegatronFusedLinearCrossEntropyFunction as fn; print(fn.__module__)" + ), + ], + check=True, + capture_output=True, + text=True, + env=env, + ) + assert result.stdout.strip() == "liger_kernel.ops.cutedsl.ops.megatron_fused_linear_cross_entropy" + + +def _tp_worker(rank, world_size, file_name, dtype): + dist.init_process_group( + backend="nccl", + init_method=f"file://{file_name}", + rank=rank, + world_size=world_size, + ) + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + tp_group = dist.group.WORLD + vocab_global = 64 + vocab_local = vocab_global // world_size + + torch.manual_seed(123) + hidden_base = torch.randn(3, 2, 65, device=device, dtype=dtype) + weight_global = torch.randn(vocab_global, 65, device=device, dtype=dtype) * 0.02 + bias_global = torch.randn(vocab_global, device=device, dtype=dtype) * 0.02 + target = torch.randint(0, vocab_global, (3, 2), device=device) + upstream = torch.randn(3, 2, device=device) + target[0, 0] = -100 + for tensor in (hidden_base, weight_global, bias_global, target, upstream): + dist.broadcast(tensor, src=0, group=tp_group) + + start = rank * vocab_local + end = start + vocab_local + hidden_cutedsl = hidden_base.clone().requires_grad_(True) + weight_local = weight_global[start:end].clone().requires_grad_(True) + bias_local = bias_global[start:end].clone().requires_grad_(True) + actual = liger_megatron_fused_linear_cross_entropy( + hidden_cutedsl, + weight_local, + target, + bias=bias_local, + tp_group=tp_group, + ) + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_global.clone().requires_grad_(True) + bias_ref = bias_global.clone().requires_grad_(True) + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + actual.backward(upstream) + reference.backward(upstream) + torch.testing.assert_close(hidden_cutedsl.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_local.grad, weight_ref.grad[start:end], atol=5e-3, rtol=5e-2) + torch.testing.assert_close(bias_local.grad, bias_ref.grad[start:end], atol=5e-3, rtol=5e-2) + dist.destroy_process_group() + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="requires at least two CUDA GPUs") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_cutedsl_megatron_flce_tp2_matches_global_reference(dtype): + with tempfile.NamedTemporaryFile() as rendezvous: + mp.spawn( + _tp_worker, + args=(2, rendezvous.name, dtype), + nprocs=2, + join=True, + ) diff --git a/test/megatron/test_cutile_fused_linear_cross_entropy.py b/test/megatron/test_cutile_fused_linear_cross_entropy.py new file mode 100644 index 000000000..853c9a3e1 --- /dev/null +++ b/test/megatron/test_cutile_fused_linear_cross_entropy.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +import os +import subprocess +import sys +import tempfile + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch.nn.functional as F + +pytest.importorskip("cuda.tile") + +from liger_kernel.ops.cutile.ops.megatron_fused_linear_cross_entropy import ( # noqa: E402 + liger_megatron_fused_linear_cross_entropy, +) + + +def _reference_loss(hidden, weight, target, bias=None, ignore_index=-100): + logits = hidden.float() @ weight.float().t() + if bias is not None: + logits = logits + bias.float() + return F.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + target.reshape(-1), + reduction="none", + ignore_index=ignore_index, + ).reshape(target.shape) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CuTile FLCE requires CUDA") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@pytest.mark.parametrize("with_bias", [False, True]) +def test_cutile_megatron_flce_tp1_matches_pytorch(dtype, with_bias): + torch.manual_seed(42) + hidden_base = torch.randn(3, 2, 65, device="cuda", dtype=dtype) + weight_base = torch.randn(33, 65, device="cuda", dtype=dtype) * 0.02 + bias_base = torch.randn(33, device="cuda", dtype=dtype) * 0.02 if with_bias else None + target = torch.randint(0, 33, (3, 2), device="cuda") + target[0, 0] = -100 + upstream = torch.randn(3, 2, device="cuda") + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_base.clone().requires_grad_(True) + bias_ref = bias_base.clone().requires_grad_(True) if bias_base is not None else None + hidden_cutile = hidden_base.clone().requires_grad_(True) + weight_cutile = weight_base.clone().requires_grad_(True) + bias_cutile = bias_base.clone().requires_grad_(True) if bias_base is not None else None + + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + actual = liger_megatron_fused_linear_cross_entropy( + hidden_cutile, + weight_cutile, + target, + bias=bias_cutile, + ) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + reference.backward(upstream) + actual.backward(upstream) + torch.testing.assert_close(hidden_cutile.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_cutile.grad, weight_ref.grad, atol=5e-3, rtol=5e-2) + if with_bias: + torch.testing.assert_close(bias_cutile.grad, bias_ref.grad, atol=5e-3, rtol=5e-2) + + +def test_cutile_megatron_flce_backend_dispatch(): + env = os.environ.copy() + env["LIGER_KERNEL_IMPL"] = "cutile" + result = subprocess.run( + [ + sys.executable, + "-c", + ( + "from liger_kernel.megatron.fused_linear_cross_entropy import " + "LigerMegatronFusedLinearCrossEntropyFunction as fn; print(fn.__module__)" + ), + ], + check=True, + capture_output=True, + text=True, + env=env, + ) + assert result.stdout.strip() == "liger_kernel.ops.cutile.ops.megatron_fused_linear_cross_entropy" + + +def _tp_worker(rank, world_size, file_name, dtype): + dist.init_process_group( + backend="nccl", + init_method=f"file://{file_name}", + rank=rank, + world_size=world_size, + ) + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + tp_group = dist.group.WORLD + vocab_global = 64 + vocab_local = vocab_global // world_size + + torch.manual_seed(123) + hidden_base = torch.randn(3, 2, 65, device=device, dtype=dtype) + weight_global = torch.randn(vocab_global, 65, device=device, dtype=dtype) * 0.02 + bias_global = torch.randn(vocab_global, device=device, dtype=dtype) * 0.02 + target = torch.randint(0, vocab_global, (3, 2), device=device) + upstream = torch.randn(3, 2, device=device) + target[0, 0] = -100 + for tensor in (hidden_base, weight_global, bias_global, target, upstream): + dist.broadcast(tensor, src=0, group=tp_group) + + start = rank * vocab_local + end = start + vocab_local + hidden_cutile = hidden_base.clone().requires_grad_(True) + weight_local = weight_global[start:end].clone().requires_grad_(True) + bias_local = bias_global[start:end].clone().requires_grad_(True) + + actual = liger_megatron_fused_linear_cross_entropy( + hidden_cutile, + weight_local, + target, + bias=bias_local, + tp_group=tp_group, + ) + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_global.clone().requires_grad_(True) + bias_ref = bias_global.clone().requires_grad_(True) + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + actual.backward(upstream) + reference.backward(upstream) + torch.testing.assert_close(hidden_cutile.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_local.grad, weight_ref.grad[start:end], atol=5e-3, rtol=5e-2) + torch.testing.assert_close(bias_local.grad, bias_ref.grad[start:end], atol=5e-3, rtol=5e-2) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires at least two CUDA GPUs", +) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_cutile_megatron_flce_tp2_matches_global_reference(dtype): + with tempfile.NamedTemporaryFile() as rendezvous: + mp.spawn( + _tp_worker, + args=(2, rendezvous.name, dtype), + nprocs=2, + join=True, + ) diff --git a/test/megatron/test_fused_linear_cross_entropy.py b/test/megatron/test_fused_linear_cross_entropy.py new file mode 100644 index 000000000..af3aafa43 --- /dev/null +++ b/test/megatron/test_fused_linear_cross_entropy.py @@ -0,0 +1,253 @@ +from __future__ import annotations + +import tempfile + +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +import torch.nn.functional as F + +from liger_kernel.megatron import LigerMegatronFusedLinearCrossEntropy +from liger_kernel.megatron import liger_megatron_fused_linear_cross_entropy_output_processor +from liger_kernel.ops.megatron_fused_linear_cross_entropy import liger_megatron_fused_linear_cross_entropy + + +def _reference_loss(hidden, weight, target, bias=None, ignore_index=-100): + logits = hidden.float() @ weight.float().t() + if bias is not None: + logits = logits + bias.float() + return F.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + target.reshape(-1), + reduction="none", + ignore_index=ignore_index, + ).reshape(target.shape) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +@pytest.mark.parametrize("shape", [(2, 3, 8, 16), (3, 2, 17, 32)]) +@pytest.mark.parametrize("with_bias", [False, True]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_megatron_flce_tp1_matches_pytorch(shape, with_bias, dtype): + s, b, h, v = shape + torch.manual_seed(42) + hidden_base = torch.randn(s, b, h, device="cuda", dtype=dtype) + weight_base = torch.randn(v, h, device="cuda", dtype=dtype) * 0.02 + bias_base = torch.randn(v, device="cuda", dtype=dtype) * 0.02 if with_bias else None + target = torch.randint(0, v, (s, b), device="cuda") + target.reshape(-1)[0] = -100 + upstream = torch.randn(s, b, device="cuda") + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_base.clone().requires_grad_(True) + bias_ref = bias_base.clone().requires_grad_(True) if bias_base is not None else None + hidden_liger = hidden_base.clone().requires_grad_(True) + weight_liger = weight_base.clone().requires_grad_(True) + bias_liger = bias_base.clone().requires_grad_(True) if bias_base is not None else None + + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + actual = liger_megatron_fused_linear_cross_entropy( + hidden_liger, + weight_liger, + target, + bias=bias_liger, + ) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + reference.backward(upstream) + actual.backward(upstream) + torch.testing.assert_close(hidden_liger.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_liger.grad, weight_ref.grad, atol=5e-3, rtol=5e-2) + if with_bias: + torch.testing.assert_close(bias_liger.grad, bias_ref.grad, atol=5e-3, rtol=5e-2) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +def test_megatron_flce_module_contract(): + module = LigerMegatronFusedLinearCrossEntropy(ignore_index=-1) + hidden = torch.randn(2, 3, 8, device="cuda", dtype=torch.bfloat16) + weight = torch.randn(16, 8, device="cuda", dtype=torch.bfloat16) * 0.02 + target = torch.randint(0, 16, (2, 3), device="cuda") + target[0, 0] = -1 + + actual = module(hidden, weight, target) + reference = _reference_loss(hidden, weight, target, ignore_index=-1) + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + assert "ignore_index=-1" in repr(module) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +def test_megatron_flce_output_processor_matches_materialized_path(): + class ColumnParallelLinear: + def __init__(self, weight, bias): + self.weight = weight + self.bias = bias + self.tp_group = None + self.gather_output = False + self.sequence_parallel = False + self.gradient_accumulation_fusion = False + self.disable_grad_reduce = False + self.explicit_expert_comm = False + self.skip_bias_add = False + + torch.manual_seed(43) + hidden = torch.randn(3, 2, 8, device="cuda", dtype=torch.bfloat16) + weight = torch.randn(16, 8, device="cuda", dtype=torch.bfloat16) * 0.02 + bias = torch.randn(16, device="cuda", dtype=torch.bfloat16) * 0.02 + labels = torch.randint(16, (2, 3), device="cuda") + output_layer = ColumnParallelLinear(weight, bias) + config = SimpleNamespace( + defer_embedding_wgrad_compute=False, + mtp_num_layers=None, + use_mup=False, + ) + + actual = liger_megatron_fused_linear_cross_entropy_output_processor( + hidden_states=hidden, + output_layer=output_layer, + output_weight=None, + labels=labels, + runtime_gather_output=None, + config=config, + ) + reference = _reference_loss(hidden, weight, labels.t().contiguous(), bias).t().contiguous() + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +def test_megatron_flce_output_processor_rejects_gathered_logits(): + output_layer = SimpleNamespace( + gather_output=True, + sequence_parallel=False, + gradient_accumulation_fusion=False, + disable_grad_reduce=False, + explicit_expert_comm=False, + skip_bias_add=False, + ) + config = SimpleNamespace( + defer_embedding_wgrad_compute=False, + mtp_num_layers=None, + use_mup=False, + ) + + with pytest.raises(RuntimeError, match="not Megatron's native ColumnParallelLinear.*gathers TP logits"): + liger_megatron_fused_linear_cross_entropy_output_processor( + hidden_states=torch.empty(1, 1, 8, device="cuda", dtype=torch.bfloat16), + output_layer=output_layer, + output_weight=None, + labels=torch.zeros(1, 1, device="cuda", dtype=torch.long), + runtime_gather_output=None, + config=config, + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_megatron_flce_triton_split_k_dx_matches_pytorch(dtype): + torch.manual_seed(44) + hidden_base = torch.randn(8, 4, 64, device="cuda", dtype=dtype) + weight_base = torch.randn(8192, 64, device="cuda", dtype=dtype) * 0.02 + target = torch.randint(0, weight_base.shape[0], hidden_base.shape[:-1], device="cuda") + upstream = torch.randn_like(target, dtype=torch.float32) + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_base.clone().requires_grad_(True) + hidden_triton = hidden_base.clone().requires_grad_(True) + weight_triton = weight_base.clone().requires_grad_(True) + + reference = _reference_loss(hidden_ref, weight_ref, target) + actual = liger_megatron_fused_linear_cross_entropy(hidden_triton, weight_triton, target) + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + + reference.backward(upstream) + actual.backward(upstream) + torch.testing.assert_close(hidden_triton.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_triton.grad, weight_ref.grad, atol=5e-3, rtol=5e-2) + + +def test_megatron_flce_rejects_cpu_inputs(): + hidden = torch.randn(2, 3, 8, dtype=torch.bfloat16) + weight = torch.randn(16, 8, dtype=torch.bfloat16) + target = torch.randint(0, 16, (2, 3)) + + with pytest.raises(RuntimeError, match="requires a CUDA GPU"): + liger_megatron_fused_linear_cross_entropy(hidden, weight, target) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Megatron FLCE requires CUDA") +def test_megatron_flce_rejects_non_long_targets(): + hidden = torch.randn(2, 3, 8, device="cuda", dtype=torch.bfloat16) + weight = torch.randn(16, 8, device="cuda", dtype=torch.bfloat16) + target = torch.zeros(2, 3, device="cuda", dtype=torch.int32) + + with pytest.raises(TypeError, match="target must have dtype torch.long"): + liger_megatron_fused_linear_cross_entropy(hidden, weight, target) + + +def _tp_worker(rank, world_size, file_name, dtype): + dist.init_process_group( + backend="nccl", + init_method=f"file://{file_name}", + rank=rank, + world_size=world_size, + ) + torch.cuda.set_device(rank) + device = torch.device("cuda", rank) + tp_group = dist.group.WORLD + s, b, h, v_global = 3, 2, 17, 32 + v_local = v_global // world_size + + torch.manual_seed(123) + hidden_base = torch.randn(s, b, h, device=device, dtype=dtype) + weight_global = torch.randn(v_global, h, device=device, dtype=dtype) * 0.02 + bias_global = torch.randn(v_global, device=device, dtype=dtype) * 0.02 + target = torch.randint(0, v_global, (s, b), device=device) + upstream = torch.randn(s, b, device=device) + target.reshape(-1)[0] = -100 + for tensor in (hidden_base, weight_global, bias_global, target, upstream): + dist.broadcast(tensor, src=0, group=tp_group) + + start = rank * v_local + end = start + v_local + hidden_liger = hidden_base.clone().requires_grad_(True) + weight_local = weight_global[start:end].clone().requires_grad_(True) + bias_local = bias_global[start:end].clone().requires_grad_(True) + + actual = liger_megatron_fused_linear_cross_entropy( + hidden_liger, + weight_local, + target, + bias=bias_local, + tp_group=tp_group, + ) + + hidden_ref = hidden_base.clone().requires_grad_(True) + weight_ref = weight_global.clone().requires_grad_(True) + bias_ref = bias_global.clone().requires_grad_(True) + reference = _reference_loss(hidden_ref, weight_ref, target, bias_ref) + + torch.testing.assert_close(actual, reference, atol=5e-3, rtol=5e-2) + actual.backward(upstream) + reference.backward(upstream) + torch.testing.assert_close(hidden_liger.grad, hidden_ref.grad, atol=5e-3, rtol=5e-2) + torch.testing.assert_close(weight_local.grad, weight_ref.grad[start:end], atol=5e-3, rtol=5e-2) + torch.testing.assert_close(bias_local.grad, bias_ref.grad[start:end], atol=5e-3, rtol=5e-2) + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < 2, + reason="requires at least two CUDA GPUs", +) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +def test_megatron_flce_tp2_matches_global_reference(dtype): + with tempfile.NamedTemporaryFile() as rendezvous: + mp.spawn( + _tp_worker, + args=(2, rendezvous.name, dtype), + nprocs=2, + join=True, + ) diff --git a/test/megatron/test_monkey_patch.py b/test/megatron/test_monkey_patch.py index 634a4abd5..ad88a078a 100644 --- a/test/megatron/test_monkey_patch.py +++ b/test/megatron/test_monkey_patch.py @@ -8,7 +8,7 @@ - patching is idempotent (calling apply twice doesn't stack wrappers) - the patch is a no-op when the kernel flag is False - missing megatron-core / missing symbol path raise helpful ``ImportError``\\s -- kernel-specific dispatch contracts (e.g. CE TP>1 raises; RMSNorm only displaces the +- kernel-specific dispatch contracts (e.g. FLCE configuration guards; RMSNorm only displaces the ``WrappedTorchNorm`` fallback, not TE / Apex) - end-to-end: the patched symbol invoked with real tensors produces correct output @@ -66,6 +66,10 @@ def _install_fake_megatron_ce( tensor_parallel = types.ModuleType("megatron.core.tensor_parallel") unfused_ce = types.ModuleType("megatron.core.tensor_parallel.cross_entropy") parallel_state = types.ModuleType("megatron.core.parallel_state") + models = sys.modules.get("megatron.core.models") or types.ModuleType("megatron.core.models") + common = types.ModuleType("megatron.core.models.common") + language_module_package = types.ModuleType("megatron.core.models.common.language_module") + language_module = types.ModuleType("megatron.core.models.common.language_module.language_module") if with_fused_symbol: @@ -73,6 +77,7 @@ def original_fused_vocab_parallel_cross_entropy(vocab_parallel_logits, target, t raise AssertionError("original megatron fused kernel called — patch failed") fused_ce.fused_vocab_parallel_cross_entropy = original_fused_vocab_parallel_cross_entropy + language_module.fused_vocab_parallel_cross_entropy = original_fused_vocab_parallel_cross_entropy if with_unfused_symbol: @@ -85,6 +90,7 @@ def original_vocab_parallel_cross_entropy( raise AssertionError("original megatron unfused kernel called — patch failed") unfused_ce.vocab_parallel_cross_entropy = original_vocab_parallel_cross_entropy + tensor_parallel.vocab_parallel_cross_entropy = original_vocab_parallel_cross_entropy parallel_state.get_tensor_model_parallel_world_size = lambda: tp_size @@ -93,12 +99,21 @@ def original_vocab_parallel_cross_entropy( sys.modules["megatron.core.tensor_parallel"] = tensor_parallel sys.modules["megatron.core.tensor_parallel.cross_entropy"] = unfused_ce sys.modules["megatron.core.parallel_state"] = parallel_state + sys.modules["megatron.core.models"] = models + sys.modules["megatron.core.models.common"] = common + sys.modules["megatron.core.models.common.language_module"] = language_module_package + sys.modules["megatron.core.models.common.language_module.language_module"] = language_module megatron_core.fusions = fusions megatron_core.tensor_parallel = tensor_parallel megatron_core.parallel_state = parallel_state + megatron_core.models = models fusions.fused_cross_entropy = fused_ce tensor_parallel.cross_entropy = unfused_ce + models.common = common + common.language_module = language_module_package + language_module_package.language_module = language_module + language_module.tensor_parallel = tensor_parallel return fused_ce, unfused_ce @@ -195,6 +210,11 @@ def _uninstall_fake_megatron(): # CE side "megatron.core.parallel_state", "megatron.core.fusions.fused_cross_entropy", + "megatron.core.models.common.language_module.language_module", + "megatron.core.models.common.language_module", + "megatron.core.models.common", + "megatron.core.models.gpt.gpt_model", + "megatron.core.models.gpt", # RMSNorm side "megatron.core.models.backends", "megatron.core.models", @@ -216,6 +236,43 @@ def _uninstall_fake_megatron(): sys.modules.pop(mod, None) +def _install_fake_megatron_gpt(with_output_processor: bool = True): + _, megatron_core = _ensure_megatron_roots() + models = sys.modules.get("megatron.core.models") or types.ModuleType("megatron.core.models") + gpt = types.ModuleType("megatron.core.models.gpt") + gpt_model = types.ModuleType("megatron.core.models.gpt.gpt_model") + + if with_output_processor: + + class GPTModel: + def _postprocess(self, labels=None, output_processor=None): + return output_processor + + else: + + class GPTModel: + def _postprocess(self, labels=None): + return labels + + gpt_model.GPTModel = GPTModel + sys.modules["megatron.core.models"] = models + sys.modules["megatron.core.models.gpt"] = gpt + sys.modules["megatron.core.models.gpt.gpt_model"] = gpt_model + megatron_core.models = models + models.gpt = gpt + gpt.gpt_model = gpt_model + return gpt_model + + +@pytest.fixture +def fake_megatron_gpt(): + gpt_model = _install_fake_megatron_gpt() + try: + yield gpt_model + finally: + _uninstall_fake_megatron() + + @pytest.fixture def fake_megatron_ce(): fused_ce, unfused_ce = _install_fake_megatron_ce(tp_size=1) @@ -399,6 +456,30 @@ def test_patch_replaces_both_fused_and_unfused_symbols_in_one_call(fake_megatron assert unfused_ce.vocab_parallel_cross_entropy.__name__ == "liger_vocab_parallel_cross_entropy" +def test_patch_rebinds_loaded_language_module_fused_consumer(fake_megatron_ce): + fused_ce, _ = fake_megatron_ce + language_module = sys.modules["megatron.core.models.common.language_module.language_module"] + original = language_module.fused_vocab_parallel_cross_entropy + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + apply_liger_kernel_to_megatron(rms_norm=False, cross_entropy=True) + + assert language_module.fused_vocab_parallel_cross_entropy is fused_ce.fused_vocab_parallel_cross_entropy + assert language_module.fused_vocab_parallel_cross_entropy is not original + + +def test_patch_rebinds_loaded_tensor_parallel_export(fake_megatron_ce): + _, unfused_ce = fake_megatron_ce + tensor_parallel = sys.modules["megatron.core.tensor_parallel"] + original = tensor_parallel.vocab_parallel_cross_entropy + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + apply_liger_kernel_to_megatron(rms_norm=False, cross_entropy=True) + + assert tensor_parallel.vocab_parallel_cross_entropy is unfused_ce.vocab_parallel_cross_entropy + assert tensor_parallel.vocab_parallel_cross_entropy is not original + + def test_patch_with_cross_entropy_false_leaves_ce_symbols_untouched(fake_megatron_ce): """Default ``cross_entropy=False`` must not touch the CE symbols even if the call runs.""" fused_ce, unfused_ce = fake_megatron_ce @@ -423,10 +504,19 @@ def test_patch_is_idempotent_for_both_symbols(fake_megatron_ce): fused_first = fused_ce.fused_vocab_parallel_cross_entropy unfused_first = unfused_ce.vocab_parallel_cross_entropy + # Model a consumer restoring its import-time binding after the definitions + # were patched; the idempotent path must repair these stale references. + language_module = sys.modules["megatron.core.models.common.language_module.language_module"] + tensor_parallel = sys.modules["megatron.core.tensor_parallel"] + language_module.fused_vocab_parallel_cross_entropy = fused_first.__wrapped__ + tensor_parallel.vocab_parallel_cross_entropy = unfused_first.__wrapped__ + apply_liger_kernel_to_megatron(rms_norm=False, cross_entropy=True) # Same identity → no stacked wrapping. assert fused_ce.fused_vocab_parallel_cross_entropy is fused_first assert unfused_ce.vocab_parallel_cross_entropy is unfused_first + assert language_module.fused_vocab_parallel_cross_entropy is fused_first + assert tensor_parallel.vocab_parallel_cross_entropy is unfused_first # __wrapped__ still references the original Megatron symbol, not the first Liger wrapper. assert fused_first.__wrapped__.__name__ == "original_fused_vocab_parallel_cross_entropy" assert unfused_first.__wrapped__.__name__ == "original_vocab_parallel_cross_entropy" @@ -920,9 +1010,11 @@ def test_import_from_root(): accidental __init__.py removals so the docs' import snippets keep working.""" try: from liger_kernel.megatron import LigerMegatronCrossEntropy # noqa: F401 + from liger_kernel.megatron import LigerMegatronFusedLinearCrossEntropy # noqa: F401 from liger_kernel.megatron import LigerMegatronRMSNorm # noqa: F401 from liger_kernel.megatron import LigerMegatronSwiGLU # noqa: F401 from liger_kernel.megatron import apply_liger_kernel_to_megatron # noqa: F401 + from liger_kernel.megatron import liger_megatron_fused_linear_cross_entropy_output_processor # noqa: F401 except Exception: pytest.fail("Importing public Megatron symbols from liger_kernel.megatron failed.") @@ -943,6 +1035,59 @@ def test_public_apply_function_has_no_ce_specific_kwargs(): ) +def test_flce_patch_injects_output_processor_for_labeled_gpt_calls(fake_megatron_gpt): + from liger_kernel.megatron import apply_liger_kernel_to_megatron + from liger_kernel.megatron import liger_megatron_fused_linear_cross_entropy_output_processor + + original = fake_megatron_gpt.GPTModel._postprocess + apply_liger_kernel_to_megatron(rms_norm=False, fused_linear_cross_entropy=True) + patched = fake_megatron_gpt.GPTModel._postprocess + + assert patched is not original + assert patched.__wrapped__ is original + assert patched(object(), labels=object()) is liger_megatron_fused_linear_cross_entropy_output_processor + assert patched(object(), labels=None) is None + + +def test_flce_patch_is_opt_in(fake_megatron_gpt): + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + original = fake_megatron_gpt.GPTModel._postprocess + apply_liger_kernel_to_megatron(rms_norm=False) + + assert fake_megatron_gpt.GPTModel._postprocess is original + + +def test_flce_patch_preserves_custom_output_processor(fake_megatron_gpt): + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + custom = object() + apply_liger_kernel_to_megatron(rms_norm=False, fused_linear_cross_entropy=True) + + assert fake_megatron_gpt.GPTModel()._postprocess(labels=object(), output_processor=custom) is custom + + +def test_flce_patch_is_idempotent(fake_megatron_gpt): + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + apply_liger_kernel_to_megatron(rms_norm=False, fused_linear_cross_entropy=True) + first = fake_megatron_gpt.GPTModel._postprocess + apply_liger_kernel_to_megatron(rms_norm=False, fused_linear_cross_entropy=True) + + assert fake_megatron_gpt.GPTModel._postprocess is first + + +def test_flce_patch_requires_megatron_output_processor_hook(): + _install_fake_megatron_gpt(with_output_processor=False) + try: + from liger_kernel.megatron import apply_liger_kernel_to_megatron + + with pytest.raises(ImportError, match="Megatron-Core 0.18 or newer"): + apply_liger_kernel_to_megatron(rms_norm=False, fused_linear_cross_entropy=True) + finally: + _uninstall_fake_megatron() + + # =========================================================================== # 5. End-to-end integration through the patched CE symbols # ===========================================================================