MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / __init__

Method __init__

codegeex/mindspore/src/utils.py:141–191  ·  view source on GitHub ↗
(self, params, config)

Source from the content-addressed store, hash-verified

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"""

Callers 3

__init__Method · 0.45
__init__Method · 0.45
__init__Method · 0.45

Calls 3

_get_pipeline_groupFunction · 0.85
encodeMethod · 0.45

Tested by

no test coverage detected