Files
project_6/enginex/ops/norm.py
EngineX b4e055e9a9 feat(enginex): CCCL-style algorithm factor replacement engine — 18 operator dispatch system
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.
2026-08-10 02:40:25 +00:00

37 lines
1.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
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)