已合并
[sync] PR-40312: Update moe token unpermute DTensor support #42584
[sync] PR-40312: Update moe token unpermute DTensor support #42584
已合并
shawnylee233创建于 7月24日
1 个文件变更+54-0
Mtorch_npu/distributed/tensor/_moe_ops.py+54-0
@@ -98,6 +98,31 @@ def npu_moe_token_unpermute_strategy(permuted_tokens, sorted_indices, probs=None
98 return strategies98 return strategies
99 99 
100 100 
101+@register_sharding(npu._npu_moe_token_unpermute.default)
102+def _npu_moe_token_unpermute_strategy(permuted_tokens, sorted_indices, probs=None, padded_mode=False,
103+ restore_shape=None):
104+ # func: _npu_moe_token_unpermute(Tensor permuted_tokens, Tensor sorted_indices, Tensor? probs=None,
105+ # bool padded_mode=False, int[]? restore_shape=None)
106+ # -> (Tensor unpermuted_tokens, Tensor permuted_tokens_for_backward)
107+ strategies = []
108+ 
109+ # all replicate strategy
110+ replicate_strategy = (
111+ [Replicate(), Replicate()], # output, saved permuted_tokens
112+ [Replicate(), Replicate(), None if probs is None else Replicate(), None, None] # input
113+ )
114+ strategies.append(replicate_strategy)
115+ 
116+ # hidden_size dim sharding strategy
117+ hidden_size_sharding_strategy = (
118+ [Shard(1), Shard(1)],
119+ [Shard(1), Replicate(), None if probs is None else Replicate(), None, None]
120+ )
121+ strategies.append(hidden_size_sharding_strategy)
122+ 
123+ return strategies
124+ 
125+ 
101@register_sharding(npu.npu_moe_token_unpermute_grad.default)126@register_sharding(npu.npu_moe_token_unpermute_grad.default)
102def npu_moe_token_unpermute_grad_strategy(permuted_tokens, grad_unpermuted_tokens, sorted_indices, probs=None,127def npu_moe_token_unpermute_grad_strategy(permuted_tokens, grad_unpermuted_tokens, sorted_indices, probs=None,
103 padded_mode=False, restore_shape=None):128 padded_mode=False, restore_shape=None):
@@ -121,3 +146,32 @@ def npu_moe_token_unpermute_grad_strategy(permuted_tokens, grad_unpermuted_token
121 strategies.append(hidden_size_sharding_strategy)146 strategies.append(hidden_size_sharding_strategy)
122 147 
123 return strategies148 return strategies
149+ 
150+ 
151+@register_sharding(npu.npu_moe_token_unpermute_grad_v2.default)
152+def npu_moe_token_unpermute_grad_v2_strategy(grad_unpermuted_tokens, sorted_indices, permuted_tokens_size_0,
153+ permuted_tokens_dtype, probs=None, padded_mode=False,
154+ restore_shape=None, permuted_tokens=None):
155+ # func: npu_moe_token_unpermute_grad_v2(Tensor grad_unpermuted_tokens, Tensor sorted_indices,
156+ # int permuted_tokens_size_0, ScalarType permuted_tokens_dtype,
157+ # Tensor? probs=None, bool padded_mode=False, int[]? restore_shape=None,
158+ # Tensor? permuted_tokens=None) -> (Tensor, Tensor)
159+ strategies = []
160+ 
161+ # all replicate strategy
162+ replicate_strategy = (
163+ [Replicate(), None if probs is None else Replicate()], # permuted_tokens grad, probs grad
164+ [Replicate(), Replicate(), None, None, None if probs is None else Replicate(), None, None,
165+ None if permuted_tokens is None else Replicate()] # input
166+ )
167+ strategies.append(replicate_strategy)
168+ 
169+ # hidden_size dim sharding strategy
170+ hidden_size_sharding_strategy = (
171+ [Shard(1), Partial()],
172+ [Shard(1), Replicate(), None, None, None if probs is None else Replicate(), None, None,
173+ None if permuted_tokens is None else Shard(1)]
174+ )
175+ strategies.append(hidden_size_sharding_strategy)
176+ 
177+ return strategies