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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
110 changes: 110 additions & 0 deletions benchmark/scripts/benchmark_mlp.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
import torch

from benchmark_model_configs import MODEL_REGISTRY
from benchmark_model_configs import build_model_config_sweep
from benchmark_model_configs import build_token_length_sweep
from benchmark_model_configs import get_benchmark_model_config
from transformers.models.llama.configuration_llama import LlamaConfig
from transformers.models.llama.modeling_llama import LlamaMLP
from utils import SingleBenchmarkRunInput
from utils import build_memory_bench_fn
from utils import build_speed_bench_fn
from utils import parse_benchmark_script_args
from utils import run_benchmarks

from liger_kernel.transformers.mlp import LigerMLP
from liger_kernel.utils import infer_device

device = infer_device()


def setup_mlp(input: SingleBenchmarkRunInput):
"""Create input tensor and SwiGLU MLP layer from benchmark config."""
cfg = input.extra_benchmark_config
if isinstance(input.x, str):
model_cfg = MODEL_REGISTRY[input.x]
seq_len = cfg["seq_len"]
hidden_size = model_cfg.hidden_size
intermediate_size = model_cfg.intermediate_size
dtype = model_cfg.dtype
else:
seq_len = input.x
hidden_size = cfg["hidden_size"]
intermediate_size = cfg["intermediate_size"]
dtype = cfg["dtype"]

llama_config = LlamaConfig(
hidden_size=hidden_size,
intermediate_size=intermediate_size,
hidden_act=cfg["hidden_act"],
)
x = torch.randn(
cfg["bsz"],
seq_len,
hidden_size,
device=device,
dtype=dtype,
requires_grad=True,
)
if input.kernel_provider == "liger":
layer = LigerMLP(config=llama_config).to(device).to(dtype)
elif input.kernel_provider == "huggingface":
layer = LlamaMLP(config=llama_config).to(device).to(dtype)
else:
raise ValueError(f"Invalid provider: {input.kernel_provider} for MLP")
return x, layer


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

if args.sweep_mode == "model_config":
common_configs = build_model_config_sweep(
kernel_name="mlp",
setup_fn=setup_mlp,
model_keys=["hidden_size", "intermediate_size", "dtype"],
probe_provider="huggingface",
extra_configs={
"bsz": 1,
"hidden_act": "silu",
},
probe_dim="T",
bt=args.bt,
overwrite=args.overwrite,
)
else:
model = get_benchmark_model_config(args.model)
probe_seq_len = 1024

common_configs = build_token_length_sweep(
kernel_name="mlp",
probe_x=probe_seq_len,
model=model,
setup_fn=setup_mlp,
model_keys=["hidden_size", "intermediate_size", "dtype"],
extra_configs={
"bsz": 1,
"hidden_act": "silu",
},
scale_dim="T",
x_label="total tokens",
probe_provider="huggingface",
overwrite=args.overwrite,
)

common_configs["kernel_providers"] = ["huggingface", "liger"]

run_benchmarks(
bench_test_fn=build_speed_bench_fn(setup_mlp),
kernel_operation_modes=["forward", "backward", "full"],
metric_name="speed",
metric_unit="ms",
**common_configs,
)
run_benchmarks(
bench_test_fn=build_memory_bench_fn(setup_mlp),
kernel_operation_modes=["full", "forward", "backward"],
metric_name="memory",
metric_unit="MB",
**common_configs,
)
1 change: 1 addition & 0 deletions src/liger_kernel/ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@
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
from liger_kernel.ops.mlp import LigerMLPFunction # noqa: F401
from liger_kernel.ops.modulated_rms_norm import LigerModulatedRMSNormFunction # noqa: F401
from liger_kernel.ops.modulated_rms_norm import modulated_rms_norm_backward # noqa: F401
from liger_kernel.ops.modulated_rms_norm import modulated_rms_norm_forward # noqa: F401
Expand Down
Loading
Loading