Skip to content
Merged

Ar rms #1290

Show file tree
Hide file tree
Changes from 2 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
6 changes: 6 additions & 0 deletions aiter/dist/communication_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,12 @@ def tensor_model_parallel_all_reduce(
return get_tp_group().all_reduce(input_, open_fp8_quant)


def tensor_model_parallel_fused_allreduce_rmsnorm(
input_: torch.Tensor, weight_: torch.Tensor, eps: float
) -> torch.Tensor:
return get_tp_group().fused_allreduce_rmsnorm(input_, weight_, eps)


def tensor_model_parallel_custom_all_gather(input_: torch.Tensor) -> torch.Tensor:
return get_tp_group().custom_all_gather(input_)

Expand Down
32 changes: 32 additions & 0 deletions aiter/dist/device_communicators/communicator_cuda.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,38 @@ def all_reduce(self, input_, ca_fp8_quant: bool = False) -> torch.Tensor:
torch.distributed.all_reduce(out, group=self.device_group)
return out

def fused_allreduce_rmsnorm(self, input_, weight_, eps) -> torch.Tensor:
n = input_.shape[-1]
can_use_fuse_ar_rms = (
n <= 16384 and input_.numel() * input_.element_size() < 8 * 1024 * 8192
)
ca_comm = self.ca_comm
if (
ca_comm is not None
and not ca_comm.disabled
and ca_comm.should_custom_ar(input_)
and can_use_fuse_ar_rms
):
out = ca_comm.custom_fused_ar_rms(input_, weight_, eps)
assert out is not None
return out
# call split kernel
ar_out = all_reduce(input_)
out = torch.empty_like(ar_out)
residual_out = torch.empty_like(ar_out)
from aiter import rmsnorm2d_fwd_with_add

rmsnorm2d_fwd_with_add(
out,
ar_out,
input_,
residual_out,
weight_,
eps,
0,
)
return out

def reduce_scatter(self, input_: torch.Tensor, dim: int = -1):
world_size = self.world_size
pynccl_comm = self.pynccl_comm
Expand Down
35 changes: 35 additions & 0 deletions aiter/dist/device_communicators/custom_all_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,6 +336,41 @@ def custom_all_gather(self, inp: torch.Tensor) -> Optional[torch.Tensor]:
else:
return self.all_gather_unreg(inp)

def fused_ar_rms(
self,
inp: torch.Tensor,
*,
out: Optional[torch.Tensor] = None,
w: torch.Tensor,
eps: float,
registered: bool = False,
):
if out is None:
out = torch.empty_like(inp)
ops.fused_allreduce_rmsnorm(
self._ptr,
inp,
out,
w,
eps,
None if registered else self.buffer,
)
return out

def custom_fused_ar_rms(
self, input: torch.Tensor, weight: torch.Tensor, eps: float
) -> Optional[torch.Tensor]:
# when custom allreduce is disabled, this will be None
if self.disabled or not self.should_custom_ar(input):
return None
if self._IS_CAPTURING:
if torch.cuda.is_current_stream_capturing():
return self.fused_ar_rms(input, w=weight, eps=eps, registered=True)
else:
return torch.empty_like(input)
else:
return self.all_reduce(input, w=weight, eps=eps, registered=False)
Comment thread
valarLip marked this conversation as resolved.
Outdated

def close(self):
if not self.disabled and self._ptr:
ops.dispose(self._ptr)
Expand Down
31 changes: 31 additions & 0 deletions aiter/dist/parallel_state.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,23 @@ def all_reduce_(
return group._all_reduce_out_place(tensor, ca_fp8_quant)


def fused_allreduce_rmsnorm_fake(
inp: torch.Tensor, w: torch.Tensor, eps: float, group_name: str
) -> torch.Tensor:
return torch.empty_like(inp)


@torch_compile_guard(gen_fake=fused_allreduce_rmsnorm_fake)
def fused_allreduce_rmsnorm_(
inp: torch.Tensor, w: torch.Tensor, eps: float, group_name: str
) -> torch.Tensor:
assert group_name in _groups, f"Group {group_name} is not found."
group = _groups[group_name]()
if group is None:
raise ValueError(f"Group {group_name} is destroyed.")
return group._fused_allreduce_rmsnorm_out_place(inp, w, eps)


if supports_custom_op():

# @torch.library.custom_op("aiter::outplace_all_gather", mutates_args=[])
Expand Down Expand Up @@ -329,6 +346,20 @@ def _all_reduce_out_place(
raise ValueError("No device communicator found")
return self.device_communicator.all_reduce(input_, ca_fp8_quant)

def fused_allreduce_rmsnorm(
self, input_: torch.Tensor, weight_: torch.Tensor, eps: float
) -> torch.Tensor:
return fused_allreduce_rmsnorm_(
input_, weight_, eps, group_name=self.unique_name
)

def _fused_allreduce_rmsnorm_out_place(
self, input_: torch.Tensor, weight_: torch.Tensor, eps: float
) -> torch.Tensor:
if self.device_communicator is None:
raise ValueError("No device communicator found")
return self.device_communicator.fused_allreduce_rmsnorm(input_, weight_, eps)

def _all_gather_out_place(self, input_: torch.Tensor) -> torch.Tensor:
ca_comm = self.device_communicator.ca_comm
assert ca_comm is not None
Expand Down
11 changes: 11 additions & 0 deletions aiter/ops/custom_all_reduce.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,17 @@ def all_gather_unreg(
) -> None: ...


@compile_ops("module_custom_all_reduce")
def fused_allreduce_rmsnorm(
_fa: int,
inp: torch.Tensor,
out: torch.Tensor,
w: torch.Tensor,
eps: float,
reg_buffer: Optional[torch.Tensor] = None,
) -> None: ...


def all_reduce_asm_fake_tensor(
inp: torch.Tensor,
ca: int,
Expand Down
Loading