(self, params, config)
| 139 | """ |
| 140 | |
| 141 | def __init__(self, params, config): |
| 142 | super(GlobalNorm, self).__init__() |
| 143 | self.norm = nn.Norm() |
| 144 | self.hyper_map = C.HyperMap() |
| 145 | self.is_pipeline = context.get_auto_parallel_context("pipeline_stages") > 1 |
| 146 | optimizer_weight_shard_size = context.get_auto_parallel_context("optimizer_weight_shard_size") |
| 147 | if self.is_pipeline: |
| 148 | if context.get_auto_parallel_context("enable_parallel_optimizer"): |
| 149 | group_size = get_group_size() // config.parallel_config.pipeline_stage |
| 150 | if optimizer_weight_shard_size > 0: |
| 151 | group_size = optimizer_weight_shard_size |
| 152 | else: |
| 153 | group_size = config.parallel_config.model_parallel |
| 154 | group_list, group_name = _get_model_parallel_group(group_size) |
| 155 | # In avoid of the group name too long |
| 156 | hashed = hashlib.md5(group_name.encode()).hexdigest()[:48] |
| 157 | print(f"Creating hash value for the group_name hash({group_name})={hashed}") |
| 158 | group_name = str(hashed) |
| 159 | create_group(group_name, group_list) |
| 160 | self.allreduce = P.AllReduce(group=group_name) |
| 161 | pipeline_group_list, pipeline_group_name = _get_pipeline_group() |
| 162 | hashed = hashlib.md5(pipeline_group_name.encode()).hexdigest()[:48] |
| 163 | print(f"Creating hash value for the group_name hash({pipeline_group_name})={hashed}") |
| 164 | pipeline_group_name = str(hashed) |
| 165 | create_group(pipeline_group_name, pipeline_group_list) |
| 166 | self.allreduce2 = P.AllReduce(group=pipeline_group_name) |
| 167 | else: |
| 168 | opt_shard_size = config.parallel_config.data_parallel |
| 169 | mp = config.parallel_config.model_parallel |
| 170 | if context.get_auto_parallel_context("enable_parallel_optimizer") and optimizer_weight_shard_size > 0: |
| 171 | opt_shard_size = optimizer_weight_shard_size |
| 172 | group_size = opt_shard_size * mp |
| 173 | world_size = get_group_size() |
| 174 | dense_repeat_num = world_size // group_size |
| 175 | layernorm_and_bias_repeat_num = world_size |
| 176 | word_embbedding_repeat_num = world_size // mp |
| 177 | position_embedding_repeat_num = world_size |
| 178 | |
| 179 | self.allreduce_group_size = () |
| 180 | for x in params: |
| 181 | if "projection.bias" not in x.name and "layernorm" not in x.name and "embedding_table" not in x.name: |
| 182 | self.allreduce_group_size = self.allreduce_group_size + (dense_repeat_num * 1.0,) |
| 183 | elif "embedding_table" not in x.name: |
| 184 | self.allreduce_group_size = self.allreduce_group_size + (layernorm_and_bias_repeat_num * 1.0,) |
| 185 | else: |
| 186 | if not config.parallel_config.vocab_emb_dp and "position_embedding.embedding_table" not in x.name \ |
| 187 | and "top_query_embedding_table" not in x.name: |
| 188 | self.allreduce_group_size = self.allreduce_group_size + \ |
| 189 | (word_embbedding_repeat_num * 1.0,) |
| 190 | else: |
| 191 | self.allreduce_group_size = self.allreduce_group_size + (position_embedding_repeat_num * 1.0,) |
| 192 | |
| 193 | def construct(self, grads): |
| 194 | """Calculate global norm construct""" |
no test coverage detected