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

Method __init__

codegeex/mindspore/src/pangu_alpha.py:41–90  ·  view source on GitHub ↗
(self, config)

Source from the content-addressed store, hash-verified

39 r"""Embedding layer of the PanGUAlpha Model"""
40
41 def __init__(self, config):
42 super(EmbeddingLayer, self).__init__()
43 # Only for the pipeline mode, the embedding needs to be row sliced.
44 dp = config.parallel_config.embedding_dp_mp_config.data_parallel
45 mp = config.parallel_config.embedding_dp_mp_config.model_parallel
46 self.word_embedding = VocabEmbedding(
47 vocab_size=(config.vocab_size // 1024 + 1) * 1024,
48 embedding_size=config.hidden_size,
49 param_init=initializer(
50 "normal",
51 [(config.vocab_size // 1024 + 1) * 1024, config.hidden_size],
52 # dtype=config.param_init_type,
53 dtype=mstype.float32,
54 ),
55 parallel_config=config.parallel_config.embedding_dp_mp_config,
56 )
57 self.word_embedding.gather.shard(((mp, 1), (dp, 1)))
58 # self.word_embedding.embedding_table.parallel_optimizer = True
59 copied_parallel_config = copy.deepcopy(config.parallel_config)
60 copied_parallel_config.vocab_emb_dp = True
61 self.position_embedding = VocabEmbedding(
62 vocab_size=config.seq_length,
63 embedding_size=config.hidden_size,
64 param_init=initializer(
65 "normal",
66 [config.seq_length, config.hidden_size],
67 # dtype=config.param_init_type,
68 dtype=mstype.float32,
69 ),
70 parallel_config=copied_parallel_config.embedding_dp_mp_config,
71 )
72 self.split = P.Split(1, 2).shard(
73 ((config.parallel_config.data_parallel, 1, 1),)
74 )
75 self.print = P.Print()
76 self.add = P.Add().shard(
77 (
78 (config.parallel_config.data_parallel, 1, 1),
79 (config.parallel_config.data_parallel, 1, 1),
80 )
81 )
82 # self.dropout = nn.Dropout(1 - config.dropout_rate)
83 self.dropout = _Dropout(1 - config.dropout_rate)
84 self.dropout.shard(((config.parallel_config.data_parallel, 1, 1),))
85 # self.dropout.dropout.shard(
86 # ((config.parallel_config.data_parallel, 1, 1),)
87 # )
88 self.is_first_iteration = True
89 self.use_past = config.use_past
90 self.batch_size = config.batch_size
91
92 def construct(
93 self, input_ids, input_position, init_reset, batch_valid_length

Callers

nothing calls this directly

Calls 1

__init__Method · 0.45

Tested by

no test coverage detected