Native and eager Triton mHC post-processing, without projection or tensor caches.
logger module-attribute
logger = init_logger(__name__)
MHCMix
Bases: CustomOp
Mix four streams and add a separately rounded branch term.
forward_npu class-attribute instance-attribute
forward_cuda
forward_cuda(
streams: Tensor,
branch_output: Tensor,
post_coefficients: Tensor,
residual_matrix: Tensor,
) -> Tensor
forward_native staticmethod
forward_native(
streams: Tensor,
branch_output: Tensor,
post_coefficients: Tensor,
residual_matrix: Tensor,
) -> Tensor
MHCPostResidual
Bases: CustomOp
Prepare post coefficients and the residual matrix with FP32 arithmetic.
forward_npu class-attribute instance-attribute
forward_cuda
forward_cuda(
post_logits: Tensor,
residual_logits: Tensor,
alpha_post: Tensor,
bias_post: Tensor,
alpha_residual: Tensor,
bias_residual: Tensor,
*,
scale: float,
iterations: int = 20,
epsilon: float = 1e-12,
out_dtype: dtype,
) -> tuple[Tensor, Tensor]
forward_native staticmethod
forward_native(
post_logits: Tensor,
residual_logits: Tensor,
alpha_post: Tensor,
bias_post: Tensor,
alpha_residual: Tensor,
bias_residual: Tensor,
*,
scale: float,
iterations: int = 20,
epsilon: float = 1e-12,
out_dtype: dtype,
) -> tuple[Tensor, Tensor]
sinkhorn_knopp
sinkhorn_knopp(
matrix_logits: Tensor, iterations: int, epsilon: float
) -> Tensor