已合并
[sync] PR-40312: Update moe token unpermute DTensor support #42530
ascend-robot创建于 26 天前
[sync] PR-40312: Update moe token unpermute DTensor support #42530
已合并
共 1 个文件变更+54-0
| @@ -98,6 +98,31 @@ def npu_moe_token_unpermute_strategy(permuted_tokens, sorted_indices, probs=None | |||
| 98 | return strategies | 98 | return strategies |
| 99 | 99 | ||
| 100 | 100 | ||
| 101 | |||
| 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 | 126 | ||
| 102 | def npu_moe_token_unpermute_grad_strategy(permuted_tokens, grad_unpermuted_tokens, sorted_indices, probs=None, | 127 | def 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 strategies | 148 | return strategies |
| 149 | |||
| 150 | |||
| 151 | |||
| 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 | ||