init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

@@ -0,0 +1,231 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
xactmask = torch.randint(0, 2, (m,), torch.bool).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
scale1=scale1_npu,
scale2=scale2_npu,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
x_active_mask=xactmask,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.int8).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.int8).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
scale1=scale1_npu,
scale2=scale2_npu,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-1024, 1024, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,229 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.bfloat16).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.bfloat16).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
weight1 = self.generate_random_tensor((e, k, n), dtype=torch.bfloat16).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2), dtype=torch.bfloat16).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32).npu()
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64).npu()
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64).npu()
probs = torch.randn(size=(m, topk), dtype=torch.float32).npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=torch.tensor([]),
bias2=torch.tensor([]),
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-1024, 1024, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,335 @@
import random
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from torch.distributed.distributed_c10d import _get_default_group
from vllm_ascend.utils import enable_custom_op
enable_custom_op()
DEVICE_OFFSET = 0
def int32_to_8x_int4_float(tensor_int32):
"""
Unpack each int32 value in the tensor into 8 signed int4 values and convert them to float32.
Logic:
1. Extract the lower 4 bits -> 0th int4
2. Shift right by 4 bits, extract the lower 4 bits -> 1st int4
...
3. Shift right by 28 bits, extract the lower 4 bits -> 7th int4
For signed int4 (Two's complement):
Binary 0000 ~ 0111 (0~7) -> float 0.0 ~ 7.0
Binary 1000 ~ 1111 (8~15) -> float -8.0 ~ -1.0
"""
# Ensure the dtype is int32 (for robustness, even if the input is already int32)
if tensor_int32.dtype != torch.int32:
tensor_int32 = tensor_int32.to(torch.int32)
original_shape = tensor_int32.shape
# 1. Create shift amounts [0, 4, 8, 12, 16, 20, 24, 28]
# Reshape to (1, 1, ..., 8) for broadcasting
shifts = torch.arange(0, 32, 4, device=tensor_int32.device).view(*([1] * len(original_shape)), -1)
# 2. Expand dimension and shift right
# unsqueeze(-1) adds a dimension -> [..., 1]
# After shifting -> [..., 8]
shifted = tensor_int32.unsqueeze(-1) >> shifts
# 3. Apply mask to keep only the lower 4 bits (0xF = 1111 binary)
# The value range here is 0 ~ 15 (unsigned view)
unpacked_unsigned = shifted & 0xF
# 4. Convert to signed int4 (-8 ~ 7)
# If value >= 8, the highest bit is 1, representing a negative number.
# In two's complement, 4-bit values 8~15 correspond to -8~-1.
# Algorithm: val = val - 16 (if val >= 8)
unpacked_signed = unpacked_unsigned.to(torch.int32) # Ensure calculation precision
mask = unpacked_signed >= 8
unpacked_signed[mask] -= 16
# 5. Convert to float32
result_float = unpacked_signed.to(torch.float32)
result_flat = result_float.flatten(start_dim=-2)
return result_flat
class TestDispatchFFNCombine:
def __init__(self, rank, world_size, port):
self.rank = rank
self.world_size = world_size
self.master_ip = "127.0.0.1"
self.port = port
def get_hcomm(self, comm_group):
hcomm_info = None
if torch.__version__ > "2.0.1":
hcomm_info = comm_group._get_backend(torch.device("npu")).get_hccl_comm_name(self.rank)
else:
hcomm_info = comm_group.get_hccl_comm_name(self.rank)
return hcomm_info
def setup_ep_tp(
self,
rank,
tp_size,
ep_size,
backend_type,
ep_ranks_list=None,
tp_ranks_list=None,
):
for i in range(tp_size):
if ep_ranks_list:
ep_ranks = ep_ranks_list[i]
else:
ep_ranks = [x + ep_size * i for x in range(ep_size)]
ep_group = dist.new_group(backend=backend_type, ranks=ep_ranks)
if rank in ep_ranks:
ep_group_tmp = ep_group
for i in range(ep_size):
if tp_ranks_list:
tp_ranks = tp_ranks_list[i]
else:
tp_ranks = [x * ep_size + i for x in range(tp_size)]
tp_group = dist.new_group(backend=backend_type, ranks=tp_ranks)
if rank in tp_ranks:
tp_group_tmp = tp_group
return ep_group_tmp, tp_group_tmp
def generate_hcom(self):
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
dist.init_process_group(
backend="hccl",
rank=self.rank,
world_size=self.world_size,
init_method=f"tcp://127.0.0.1:{self.port}",
)
ep_size = 0
tp_size = self.world_size
hcomm_info_dist = {
"default_pg_info": None,
"ep_hcomm_info": None,
"group_ep": None,
"tp_hcomm_info": None,
"group_tp": None,
}
if ep_size and tp_size:
group_ep, group_tp = self.setup_ep_tp(self.rank, tp_size, ep_size, "hccl", None, None)
hcomm_info_dist["ep_hcomm_info"] = self.get_hcomm(group_ep)
hcomm_info_dist["tp_hcomm_info"] = self.get_hcomm(group_tp)
hcomm_info_dist["group_ep"] = group_ep
hcomm_info_dist["group_tp"] = group_tp
else:
if dist.is_available():
default_pg = _get_default_group()
hcomm_info_dist["default_pg_info"] = self.get_hcomm(default_pg)
hcomm_info = hcomm_info_dist["default_pg_info"]
self.hcomm_info = hcomm_info
def run_tensor_list(self) -> bool:
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
active_num = m // 8
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16)
weight1 = self.generate_random_tensor((e, k, n // 8), dtype=torch.int32).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2 // 8), dtype=torch.int32).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
bias1 = int32_to_8x_int4_float(weight1.cpu())
bias1_npu = bias1.sum(dim=-1).npu() # shape: [e, n]
bias2 = int32_to_8x_int4_float(weight2.cpu())
bias2_npu = bias2.sum(dim=-1).npu() # shape: [e, n2]
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32)
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64)
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64)
probs = torch.randn(size=(m, topk), dtype=torch.float32)
x_active_mask = torch.cat(
[
torch.ones(active_num, dtype=torch.bool),
torch.zeros(m - active_num, dtype=torch.bool),
]
)
x[active_num:, :] = 0
expert_idx[active_num:, :] = torch.arange(topk, dtype=torch.int32)
x = x.npu()
expert_idx = expert_idx.npu()
scale1 = scale1.npu()
scale2 = scale2.npu()
probs = probs.npu()
x_active_mask = x_active_mask.npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
bias1_list = []
bias2_list = []
for i in range(e):
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1[i].npu(), 29))
scale1_npu.append(scale1[i].npu())
bias1_list.append(bias1_npu[i])
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2[i].npu(), 29))
scale2_npu.append(scale2[i].npu())
bias2_list.append(bias2_npu[i])
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=bias1_list,
bias2=bias2_list,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
x_active_mask=x_active_mask,
)
return True
def run_normal(self) -> bool:
torch_npu.npu.set_device(DEVICE_OFFSET + self.rank)
m = 64
k = 1024
n = 1024
topk = 8
e = 8
k2 = n // 2
n2 = k
active_num = m // 2
torch_npu.npu.config.allow_internal_format = True
x = self.generate_random_tensor((m, k), dtype=torch.bfloat16)
weight1 = self.generate_random_tensor((e, k, n // 8), dtype=torch.int32).npu()
weight1 = torch_npu.npu_format_cast(weight1, 29)
weight2 = self.generate_random_tensor((e, k2, n2 // 8), dtype=torch.int32).npu()
weight2 = torch_npu.npu_format_cast(weight2, 29)
bias1 = int32_to_8x_int4_float(weight1.cpu())
bias1_npu = bias1.sum(dim=-1).npu() # shape: [e, n]
bias2 = int32_to_8x_int4_float(weight2.cpu())
bias2_npu = bias2.sum(dim=-1).npu() # shape: [e, n2]
expert_idx = torch.randint(0, self.world_size * e, (m, topk), dtype=torch.int32)
scale1 = torch.randint(0, 1, (e, n), dtype=torch.int64)
scale2 = torch.randint(0, 1, (e, n2), dtype=torch.int64)
probs = torch.randn(size=(m, topk), dtype=torch.float32)
x_active_mask = torch.cat(
[
torch.ones(active_num, dtype=torch.bool),
torch.zeros(m - active_num, dtype=torch.bool),
]
)
x[active_num:, :] = 0
expert_idx[active_num:, :] = torch.arange(topk, dtype=torch.int32)
x = x.npu()
expert_idx = expert_idx.npu()
scale1 = scale1.npu()
scale2 = scale2.npu()
probs = probs.npu()
x_active_mask = x_active_mask.npu()
weight1_nz_npu = []
weight2_nz_npu = []
scale1_npu = []
scale2_npu = []
bias1_list = []
bias2_list = []
weight1_nz_npu.append(torch_npu.npu_format_cast(weight1.npu(), 29))
scale1_npu.append(scale1.npu())
bias1_list.append(bias1_npu)
weight2_nz_npu.append(torch_npu.npu_format_cast(weight2.npu(), 29))
scale2_npu.append(scale2.npu())
bias2_list.append(bias2_npu)
out = self.generate_random_tensor((m, k), dtype=torch.bfloat16).npu()
expert_token_nums = self.generate_random_tensor((1, e), dtype=torch.int32).npu()
torch.ops._C_ascend.dispatch_ffn_combine(
x=x,
weight1=weight1_nz_npu,
weight2=weight2_nz_npu,
expert_idx=expert_idx,
scale1=scale1_npu,
scale2=scale2_npu,
bias1=bias1_list,
bias2=bias2_list,
probs=probs,
group=self.hcomm_info,
max_output_size=512,
out=out,
expert_token_nums=expert_token_nums,
x_active_mask=x_active_mask,
)
return True
def generate_random_tensor(self, size, dtype):
if dtype in [torch.float16, torch.bfloat16, torch.float32]:
return torch.randn(size=size, dtype=dtype)
elif dtype is torch.int8:
return torch.randint(-16, 16, size=size, dtype=dtype)
elif dtype is torch.int32:
return torch.randint(-127, 127, size=size, dtype=dtype)
else:
raise ValueError(f"Invalid dtype: {dtype}")
def worker(rank: int, world_size: int, port: int, q: mp.SimpleQueue):
op = TestDispatchFFNCombine(rank, world_size, port)
op.generate_hcom()
out1 = op.run_tensor_list()
q.put(out1)
out2 = op.run_normal()
q.put(out2)
@torch.inference_mode()
def test_dispatch_ffn_combine_kernel():
world_size = 2
mp.set_start_method("fork", force=True)
q = mp.SimpleQueue()
p_list = []
port = 29501 + random.randint(0, 10000)
for rank in range(world_size):
p = mp.Process(target=worker, args=(rank, world_size, port, q))
p.start()
p_list.append(p)
results = [q.get() for _ in range(world_size)]
for p in p_list:
p.join()
assert all(results)

View File

@@ -0,0 +1,576 @@
import gc
import os
import sys
import time
from pathlib import Path
import npugraph_ex as nge
import numpy as np
import torch
import torch.distributed as dist
import torch.multiprocessing as mp
import torch_npu
from vllm_ascend.utils import enable_custom_op
torch.manual_seed(42)
torch_npu.npu.config.allow_internal_format = True
enable_custom_op()
LOG_NAME = "dispatch_gmm_combine_decode_test_logs"
BASE_KWARGS = {
"batch_size": 64,
"token_hidden_size": 7168,
"moe_intermediate_size": 2048,
"ep_world_size": 16,
"moe_expert_num": 64,
"shared_expert_rank_num": 0,
"top_k": 8,
"test_bfloat16": True,
"enable_dynamic_bs": False,
"test_graph": False,
"with_mc2_mask": False,
"dynamic_eplb": False,
"w8a8_dynamic": True,
"is_nz": True,
}
def redirect_output(log_file_path):
log_path = Path(LOG_NAME) / log_file_path
log_path.parent.mkdir(parents=True, exist_ok=True)
f = open(LOG_NAME + "/" + log_file_path, "w") # noqa: SIM115
os.dup2(f.fileno(), sys.stdout.fileno())
os.dup2(f.fileno(), sys.stderr.fileno())
return f
def permute_weight(w: torch.Tensor, tile_n):
*dims, n = w.shape
order = list(range(len(dims))) + [-2, -3, -1]
return w.reshape(*dims, 2, n // tile_n, tile_n // 2).permute(order).reshape(*dims, n).contiguous()
def output_to_file(rank_id):
return rank_id > 0
class DecodeMoeOps(torch.nn.Module):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__()
if w8a8_dynamic:
assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, (
"gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic"
)
else:
assert gmm1_weight_scale is None and gmm2_weight_scale is None, (
"gmm1_weight_scale and gmm2_weight_scale must be None for w8a8_dynamic"
)
self.ep_hcomm_info = ep_hcomm_info
self.batch_size = batch_size
self.token_hidden_size = token_hidden_size
self.moe_intermediate_size = moe_intermediate_size
self.ep_world_size = ep_world_size
self.moe_expert_num = moe_expert_num
self.global_rank_id = global_rank_id
self.shared_expert_rank_num = shared_expert_rank_num
is_shared_expert = global_rank_id < shared_expert_rank_num
moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num)
self.local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank
self.ep_recv_count_size = self.local_expert_num * ep_world_size
self.dynamic_eplb = dynamic_eplb
self.w8a8_dynamic = w8a8_dynamic
self.is_nz = is_nz
self.gmm1_weight = torch.empty([self.local_expert_num, self.token_hidden_size, self.moe_intermediate_size * 2])
self.gmm2_weight = torch.empty([self.local_expert_num, self.moe_intermediate_size, self.token_hidden_size])
if self.w8a8_dynamic:
self.gmm1_weight_scale = torch.empty([self.local_expert_num, self.moe_intermediate_size * 2])
self.gmm2_weight_scale = torch.empty([self.local_expert_num, self.token_hidden_size])
else:
self.gmm1_weight_scale = None
self.gmm2_weight_scale = None
self.gmm1_weight_scale_fp32 = None
self.gmm2_weight_scale_fp32 = None
self._process_weights_after_loading(gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale)
def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale):
if self.w8a8_dynamic:
gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ)
gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ)
self.gmm1_weight = torch.nn.Parameter(gmm1_weight, requires_grad=False)
self.gmm2_weight = torch.nn.Parameter(gmm2_weight, requires_grad=False)
if self.w8a8_dynamic:
self.gmm1_weight_scale = torch.nn.Parameter(gmm1_weight_scale, requires_grad=False)
self.gmm2_weight_scale = torch.nn.Parameter(gmm2_weight_scale, requires_grad=False)
self.gmm1_weight_scale_fp32 = torch.nn.Parameter(gmm1_weight_scale.float(), requires_grad=False)
self.gmm2_weight_scale_fp32 = torch.nn.Parameter(gmm2_weight_scale.float(), requires_grad=False)
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
raise NotImplementedError("To be implemented in subclass")
def forward(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
return self._apply_ops(x, expert_ids, smooth_scales, expert_scales, x_active_mask)
class SmallOps(DecodeMoeOps):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__(
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
dynamic_eplb,
w8a8_dynamic,
is_nz,
)
self.tp_hcomm_info = ""
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
outputs = torch_npu.npu_moe_distribute_dispatch_v2(
x=x,
expert_ids=expert_ids,
expert_scales=expert_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_world_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
group_tp=self.tp_hcomm_info,
tp_world_size=1,
tp_rank_id=0,
expert_shard_type=0,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
quant_mode=2 if self.w8a8_dynamic else 0,
global_bs=self.batch_size * self.ep_world_size,
expert_token_nums_type=1, # 0代表前缀和,1代表各自数量
)
(
expand_x,
dynamic_scales,
assist_info_for_combine,
expert_token_nums,
ep_send_counts,
tp_send_counts,
expand_scales,
) = outputs
output_dtype = x.dtype
y1_int32 = torch_npu.npu_grouped_matmul(
x=[expand_x],
weight=[self.gmm1_weight],
split_item=3,
group_list_type=1, # 默认为0,代表前缀和形式
group_type=0, # 0代表m轴分组
group_list=expert_token_nums,
output_dtype=torch.int32 if self.w8a8_dynamic else output_dtype,
)[0]
y1_scale = None
if self.w8a8_dynamic:
y1, y1_scale = torch_npu.npu_dequant_swiglu_quant(
x=y1_int32,
weight_scale=self.gmm1_weight_scale.to(torch.float32),
activation_scale=dynamic_scales,
bias=None,
quant_scale=None,
quant_offset=None,
group_index=expert_token_nums,
activate_left=True,
quant_mode=1,
)
else:
y1 = torch_npu.npu_swiglu(y1_int32)
y2 = torch_npu.npu_grouped_matmul(
x=[y1],
weight=[self.gmm2_weight],
scale=[self.gmm2_weight_scale] if self.w8a8_dynamic else None,
per_token_scale=[y1_scale] if self.w8a8_dynamic else None,
split_item=2,
group_list_type=1,
group_type=0,
group_list=expert_token_nums,
output_dtype=output_dtype,
)[0]
combine_output = torch_npu.npu_moe_distribute_combine_v2(
expand_x=y2,
expert_ids=expert_ids,
assist_info_for_combine=assist_info_for_combine,
ep_send_counts=ep_send_counts,
expert_scales=expert_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_world_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
tp_send_counts=tp_send_counts,
expand_scales=expand_scales,
group_tp=self.tp_hcomm_info,
tp_world_size=1,
tp_rank_id=0,
expert_shard_type=0,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
global_bs=self.batch_size * self.ep_world_size,
)
return (combine_output, expert_token_nums)
class FusionOp(DecodeMoeOps):
def __init__(
self,
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
super().__init__(
gmm1_weight,
gmm1_weight_scale,
gmm2_weight,
gmm2_weight_scale,
ep_hcomm_info,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
dynamic_eplb,
w8a8_dynamic,
is_nz,
)
def _apply_ops(self, x, expert_ids, smooth_scales, expert_scales, x_active_mask):
smooth_scales = torch.zeros(128 * 1024 * 1024).npu()
output = torch.ops._C_ascend.dispatch_gmm_combine_decode(
x=x,
expert_ids=expert_ids,
gmm1_permuted_weight=self.gmm1_weight,
gmm1_permuted_weight_scale=self.gmm1_weight_scale_fp32,
gmm2_weight=self.gmm2_weight,
gmm2_weight_scale=self.gmm2_weight_scale_fp32,
expert_scales=expert_scales,
expert_smooth_scales=smooth_scales,
x_active_mask=x_active_mask,
group_ep=self.ep_hcomm_info,
ep_rank_size=self.ep_world_size,
ep_rank_id=self.global_rank_id,
moe_expert_num=self.moe_expert_num,
shared_expert_num=1,
shared_expert_rank_num=self.shared_expert_rank_num,
quant_mode=0,
global_bs=self.batch_size * self.ep_world_size,
)
return output
def _process_weights_after_loading(self, gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale):
if self.is_nz:
gmm1_weight = torch_npu.npu_format_cast(gmm1_weight, torch_npu.Format.FRACTAL_NZ)
gmm2_weight = torch_npu.npu_format_cast(gmm2_weight, torch_npu.Format.FRACTAL_NZ)
if self.dynamic_eplb:
self.gmm1_weight = [weight.clone() for weight in gmm1_weight.unbind(dim=0)]
self.gmm2_weight = [weight.clone() for weight in gmm2_weight.unbind(dim=0)]
if self.w8a8_dynamic:
self.gmm1_weight_scale_fp32 = [weight.clone() for weight in gmm1_weight_scale.unbind(dim=0)]
self.gmm2_weight_scale_fp32 = [weight.clone() for weight in gmm2_weight_scale.unbind(dim=0)]
else:
self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)]
self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)]
else:
self.gmm1_weight = [gmm1_weight.clone()]
self.gmm2_weight = [gmm2_weight.clone()]
if self.w8a8_dynamic:
self.gmm1_weight_scale_fp32 = [gmm1_weight_scale.clone()]
self.gmm2_weight_scale_fp32 = [gmm2_weight_scale.clone()]
else:
self.gmm1_weight_scale_fp32 = [torch.ones(1).npu().to(gmm1_weight.dtype)]
self.gmm2_weight_scale_fp32 = [torch.ones(1).npu().to(gmm2_weight.dtype)]
def generate_datas(
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num=0,
top_k=8,
test_bfloat16=True,
enable_dynamic_bs=False,
with_mc2_mask=False,
w8a8_dynamic=True,
):
is_shared_expert = global_rank_id < shared_expert_rank_num
moe_expert_num_per_rank = moe_expert_num // (ep_world_size - shared_expert_rank_num)
actual_bs = int(
torch.randint(2 if with_mc2_mask else 1, batch_size, [1]).item() if enable_dynamic_bs else batch_size
)
local_expert_num = 1 if is_shared_expert else moe_expert_num_per_rank
gmm1_input_dim = token_hidden_size
gmm1_output_dim = moe_intermediate_size * 2
gmm2_input_dim = moe_intermediate_size
gmm2_output_dim = token_hidden_size
x = torch.rand([actual_bs, token_hidden_size]) * 0.5 - 0.5
expert_ids = (
torch.arange(global_rank_id * batch_size * top_k, global_rank_id * batch_size * top_k + actual_bs * top_k)
.to(torch.int32)
.view(actual_bs, top_k)
)
expert_ids = expert_ids % moe_expert_num
gmm1_weight_scale = None
gmm2_weight_scale = None
if w8a8_dynamic:
if is_shared_expert:
gmm1_weight = torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8) * 4
gmm2_weight = torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8) * 4
gmm1_weight[:, :, ::2] = gmm1_weight[:, :, ::2] * -1
gmm2_weight[:, :, ::2] = gmm2_weight[:, :, ::2] * -1
gmm1_weight_scale = torch.ones([local_expert_num, gmm1_output_dim]) * 0.0015
gmm2_weight_scale = torch.ones([local_expert_num, gmm2_output_dim]) * 0.0015
else:
gmm1_weight = torch.randint(-16, 16, [local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(torch.int8)
gmm2_weight = torch.randint(-16, 16, [local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(torch.int8)
gmm1_weight_scale = torch.rand([local_expert_num, gmm1_output_dim]) * 0.003 + 0.0015
gmm2_weight_scale = torch.rand([local_expert_num, gmm2_output_dim]) * 0.003 + 0.0015
else:
if is_shared_expert:
gmm1_weight = (
torch.ones([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.5
)
gmm2_weight = (
torch.ones([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.5
)
else:
gmm1_weight = (
torch.rand([local_expert_num, gmm1_input_dim, gmm1_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.25
)
gmm2_weight = (
torch.rand([local_expert_num, gmm2_input_dim, gmm2_output_dim]).to(
torch.bfloat16 if test_bfloat16 else torch.float16
)
* 0.25
)
gmm1_weight[:, ::2, :] = gmm1_weight[:, ::2, :] * -1
gmm2_weight[:, ::2, :] = gmm2_weight[:, ::2, :] * -1
expert_scales = torch.rand(actual_bs, top_k)
if test_bfloat16:
x = x.bfloat16()
if w8a8_dynamic:
assert gmm1_weight_scale is not None and gmm2_weight_scale is not None, (
"gmm1_weight_scale and gmm2_weight_scale must be provided for w8a8_dynamic"
)
gmm1_weight_scale = gmm1_weight_scale.bfloat16()
gmm2_weight_scale = gmm2_weight_scale.bfloat16()
else:
x = x.half()
smooth_sales = None
x_active_mask = None
valid_token_num = actual_bs
if with_mc2_mask:
valid_token_num = int(torch.randint(1, actual_bs, [1]).item())
x_active_mask = torch.cat((torch.ones(valid_token_num), torch.zeros(actual_bs - valid_token_num))).bool()
return (
(x, expert_ids, smooth_sales, expert_scales, x_active_mask),
(gmm1_weight, gmm1_weight_scale, gmm2_weight, gmm2_weight_scale),
actual_bs,
valid_token_num,
)
def run_once(
local_rank_id,
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
shared_expert_rank_num=0,
top_k=8,
test_bfloat16=True,
enable_dynamic_bs=False,
test_graph=False,
with_mc2_mask=False,
dynamic_eplb=False,
w8a8_dynamic=True,
is_nz=True,
):
log_file = redirect_output(f"local_rank_{local_rank_id}.log") if output_to_file(local_rank_id) else None
global_rank_id = local_rank_id # 单机
device_id = local_rank_id % 16
torch_npu.npu.set_device(device_id)
# 初始化分布式环境
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = "29500" # 端口号随意
dist.init_process_group(backend="hccl", rank=local_rank_id, world_size=ep_world_size)
ep_ranks_list = list(np.arange(0, ep_world_size))
ep_group = dist.new_group(backend="hccl", ranks=ep_ranks_list)
ep_group_small = dist.new_group(backend="hccl", ranks=ep_ranks_list)
ep_hcomm_info_fused = ep_group._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id)
ep_hcomm_info_small = ep_group_small._get_backend(torch.device("npu")).get_hccl_comm_name(local_rank_id)
torch_npu.npu.synchronize(device_id)
parameter = (
batch_size,
token_hidden_size,
moe_intermediate_size,
ep_world_size,
moe_expert_num,
global_rank_id,
shared_expert_rank_num,
)
input_datas, weight_datas, actual_bs, valid_token_num = generate_datas(
*parameter, top_k, test_bfloat16, enable_dynamic_bs, with_mc2_mask, w8a8_dynamic
)
input_datas = [data.npu() if data is not None else None for data in input_datas]
weight_datas = [data.npu() if data is not None else None for data in weight_datas]
small_ops = SmallOps(*weight_datas, ep_hcomm_info_small, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() # type: ignore
fused_ops = FusionOp(*weight_datas, ep_hcomm_info_fused, *parameter, dynamic_eplb, w8a8_dynamic, is_nz).npu() # type: ignore
if test_graph:
config = nge.CompilerConfig()
npu_backend = nge.get_npu_backend(compiler_config=config)
fused_ops = torch.compile(fused_ops, backend=npu_backend)
# test performance
start_time = time.perf_counter()
for _ in range(100):
small_op_token_output, small_op_count_output = small_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
end_time = time.perf_counter()
elapsed_time = end_time - start_time
elapsed_time_us = elapsed_time * 1000000
print(f"rank-{global_rank_id} small {elapsed_time_us} us")
start_time = time.perf_counter()
for _ in range(100):
fused_op_token_output, fused_op_count_output = fused_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
end_time = time.perf_counter()
elapsed_time = end_time - start_time
elapsed_time_us = elapsed_time * 1000000
print(f"rank-{global_rank_id} fused {elapsed_time_us} us")
small_op_token_output, small_op_count_output = small_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
print(f"rank-{global_rank_id} Small op End")
fused_op_token_output, fused_op_count_output = fused_ops(*input_datas)
torch_npu.npu.synchronize(device_id)
print(f"rank-{global_rank_id} Fused op End")
dist.destroy_process_group()
if log_file is not None:
log_file.close()
try:
torch.testing.assert_close(
small_op_token_output[0:valid_token_num].cpu(),
fused_op_token_output[0:valid_token_num].cpu(),
atol=2.0,
rtol=0.02,
)
torch.testing.assert_close(small_op_count_output.cpu(), fused_op_count_output.cpu())
except Exception as e:
print(f"rank-{global_rank_id} Assert close Failed: {e}")
else:
print(f"rank-{global_rank_id} Assert close Pass")
gc.collect()
torch.npu.empty_cache()
torch.npu.reset_peak_memory_stats()
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_base():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["batch_size"] = 32
custom_kwargs["ep_world_size"] = 8
custom_kwargs["moe_expert_num"] = 32
custom_kwargs["w8a8_dynamic"] = False
custom_kwargs["is_nz"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
print(f"{custom_kwargs=}")
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
print(f"{custom_kwargs=}")
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_with_mc2_mask():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["with_mc2_mask"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
@torch.inference_mode()
def test_dispatch_gmm_combine_decode_dynamic_eplb():
custom_kwargs = BASE_KWARGS.copy()
custom_kwargs["dynamic_eplb"] = True
ep_world_size = custom_kwargs["ep_world_size"]
custom_args = tuple(custom_kwargs.values())
mp.spawn(run_once, args=custom_args, nprocs=ep_world_size, join=True)
if __name__ == "__main__":
test_dispatch_gmm_combine_decode_base()