fix post_residual_addition more generally (#17286)
This commit is contained in:
@@ -140,6 +140,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
out, _, residual_out = torch_npu.npu_add_rms_norm(
|
out, _, residual_out = torch_npu.npu_add_rms_norm(
|
||||||
residual, x, self.weight.data, self.variance_epsilon
|
residual, x, self.weight.data, self.variance_epsilon
|
||||||
)
|
)
|
||||||
@@ -155,6 +157,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if residual is not None:
|
if residual is not None:
|
||||||
residual_out = torch.empty_like(x)
|
residual_out = torch.empty_like(x)
|
||||||
output = torch.empty_like(x)
|
output = torch.empty_like(x)
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
fused_add_rms_norm(
|
fused_add_rms_norm(
|
||||||
output,
|
output,
|
||||||
x,
|
x,
|
||||||
@@ -178,6 +182,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if residual is not None:
|
if residual is not None:
|
||||||
out = torch.empty_like(x)
|
out = torch.empty_like(x)
|
||||||
residual_out = torch.empty_like(x)
|
residual_out = torch.empty_like(x)
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
fused_add_rms_norm(
|
fused_add_rms_norm(
|
||||||
out, x, residual_out, residual, self.weight.data, self.variance_epsilon
|
out, x, residual_out, residual, self.weight.data, self.variance_epsilon
|
||||||
)
|
)
|
||||||
@@ -244,6 +250,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if _is_cpu_amx_available:
|
if _is_cpu_amx_available:
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
torch.ops.sgl_kernel.fused_add_rmsnorm_cpu(
|
torch.ops.sgl_kernel.fused_add_rmsnorm_cpu(
|
||||||
x, residual, self.weight.data, self.variance_epsilon
|
x, residual, self.weight.data, self.variance_epsilon
|
||||||
)
|
)
|
||||||
@@ -263,6 +271,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if self.variance_size_override is not None:
|
if self.variance_size_override is not None:
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
return self.forward_native(x, residual, post_residual_addition)
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon)
|
fused_add_rmsnorm(x, residual, self.weight.data, self.variance_epsilon)
|
||||||
return x, residual
|
return x, residual
|
||||||
out = rmsnorm(x, self.weight.data, self.variance_epsilon)
|
out = rmsnorm(x, self.weight.data, self.variance_epsilon)
|
||||||
@@ -284,6 +294,8 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if get_tensor_model_parallel_world_size() > 1:
|
if get_tensor_model_parallel_world_size() > 1:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
fused_result = flashinfer_allreduce_residual_rmsnorm(
|
fused_result = flashinfer_allreduce_residual_rmsnorm(
|
||||||
input_tensor=x,
|
input_tensor=x,
|
||||||
residual=residual,
|
residual=residual,
|
||||||
@@ -389,6 +401,8 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
gemma_fused_add_rmsnorm(
|
gemma_fused_add_rmsnorm(
|
||||||
x, residual, self.weight.data, self.variance_epsilon
|
x, residual, self.weight.data, self.variance_epsilon
|
||||||
)
|
)
|
||||||
@@ -404,6 +418,8 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
orig_dtype = x.dtype
|
orig_dtype = x.dtype
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
x = x + residual
|
x = x + residual
|
||||||
residual = x
|
residual = x
|
||||||
|
|
||||||
@@ -430,6 +446,8 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if _is_cpu_amx_available:
|
if _is_cpu_amx_available:
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu(
|
torch.ops.sgl_kernel.gemma_fused_add_rmsnorm_cpu(
|
||||||
x, residual, self.weight.data, self.variance_epsilon
|
x, residual, self.weight.data, self.variance_epsilon
|
||||||
)
|
)
|
||||||
@@ -446,6 +464,8 @@ class GemmaRMSNorm(MultiPlatformOp):
|
|||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
x = x + residual
|
x = x + residual
|
||||||
residual = x
|
residual = x
|
||||||
|
|
||||||
|
|||||||
@@ -50,6 +50,8 @@ def npu_wrapper_rmsnorm_forward(func):
|
|||||||
if not x.is_contiguous():
|
if not x.is_contiguous():
|
||||||
x = x.contiguous()
|
x = x.contiguous()
|
||||||
if residual is not None:
|
if residual is not None:
|
||||||
|
if post_residual_addition is not None:
|
||||||
|
residual = residual + post_residual_addition
|
||||||
out, residual_out = add_rmsnorm_bias(
|
out, residual_out = add_rmsnorm_bias(
|
||||||
x,
|
x,
|
||||||
residual,
|
residual,
|
||||||
|
|||||||
Reference in New Issue
Block a user