EngineX replaces the missing corex_gdn/corex_moe/corex_fa2 operator chain
that Sub168 has but our BI-V100 image lacks.
Architecture (mirrors CCCL dispatch/tuning/kernel three-layer system):
Registry (policy_selector) → three-tier dispatch:
Tier 1: Native .so via dlopen (libcorex_gdn.so, libixattn.so)
Tier 2: ixformer Python ops (vendor-provided)
Tier 3: PyTorch fallback (always available)
Critical fixes vs comp 168 docker log:
- moe_topk_softmax: replacement for missing ixformer op
- gdn_prefill: NaN-stable chunked impl (chunk_size=16)
- gdn_decode: state clamp prevents NaN accumulation
18 operators, all tests pass.
37 lines
1.1 KiB
Python
37 lines
1.1 KiB
Python
"""
|
||
EngineX norm operators.
|
||
|
||
RMSNorm is called 128 times per forward pass (pre-attn + post-attn × 64 layers).
|
||
fused_add_rms_norm fuses residual addition with normalization.
|
||
|
||
ixformer provides both natively. Fallbacks for environments without ixformer.
|
||
"""
|
||
|
||
import torch
|
||
|
||
|
||
def rms_norm_pytorch(
|
||
input: torch.Tensor,
|
||
weight: torch.Tensor,
|
||
output: torch.Tensor,
|
||
epsilon: float = 1e-6,
|
||
) -> None:
|
||
"""RMSNorm: output = (input / rms(input)) * weight"""
|
||
variance = input.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
||
normed = input * torch.rsqrt(variance + epsilon)
|
||
output.copy_(normed * weight)
|
||
|
||
|
||
def fused_add_rms_norm_pytorch(
|
||
input: torch.Tensor,
|
||
residual: torch.Tensor,
|
||
weight: torch.Tensor,
|
||
epsilon: float = 1e-6,
|
||
) -> None:
|
||
"""Fused: input = RMSNorm(input + residual); residual = input + residual"""
|
||
# In-place: residual += input, then normalize
|
||
residual.add_(input)
|
||
variance = residual.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
||
normed = residual * torch.rsqrt(variance + epsilon)
|
||
input.copy_(normed * weight)
|