已合并
[sync] PR-40312: Update moe token unpermute DTensor support #42530
[sync] PR-40312: Update moe token unpermute DTensor support #42530
已合并
ascend-robot创建于 26 天前
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)
102def _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)
152def 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