MCPcopy Create free account
hub / github.com/THUDM/GLM / __init__

Method __init__

model/modeling_glm.py:47–95  ·  view source on GitHub ↗
(self,
                 num_layers,
                 vocab_size,
                 hidden_size,
                 num_attention_heads,
                 embedding_dropout_prob,
                 attention_dropout_prob,
                 output_dropout_prob,
                 max_sequence_length,
                 max_memory_length,
                 checkpoint_activations,
                 checkpoint_num_layers=1,
                 parallel_output=True,
                 relative_encoding=False,
                 block_position_encoding=False,
                 output_predict=True,
                 spell_length=None,
                 spell_func='lstm',
                 attention_scale=1.0,
                 )

Source from the content-addressed store, hash-verified

45 """
46
47 def __init__(self,
48 num_layers,
49 vocab_size,
50 hidden_size,
51 num_attention_heads,
52 embedding_dropout_prob,
53 attention_dropout_prob,
54 output_dropout_prob,
55 max_sequence_length,
56 max_memory_length,
57 checkpoint_activations,
58 checkpoint_num_layers=1,
59 parallel_output=True,
60 relative_encoding=False,
61 block_position_encoding=False,
62 output_predict=True,
63 spell_length=None,
64 spell_func='lstm',
65 attention_scale=1.0,
66 ):
67
68 super(GLMModel, self).__init__()
69
70 self.parallel_output = parallel_output
71 self.output_predict = output_predict
72 self.hidden_size = hidden_size
73
74 init_method = init_method_normal(std=0.02)
75
76 # Word embeddings (parallel).
77 self.word_embeddings = mpu.VocabParallelEmbedding(
78 vocab_size, hidden_size, init_method=init_method)
79
80 # Transformer
81 self.transformer = mpu.GPT2ParallelTransformer(num_layers,
82 hidden_size,
83 num_attention_heads,
84 max_sequence_length,
85 max_memory_length,
86 embedding_dropout_prob,
87 attention_dropout_prob,
88 output_dropout_prob,
89 checkpoint_activations,
90 checkpoint_num_layers,
91 attention_scale=attention_scale,
92 relative_encoding=relative_encoding,
93 block_position_encoding=block_position_encoding)
94 if spell_length is not None:
95 self.prompt_spell = PromptSpell(spell_length, self.hidden_size, spell_func)
96
97 def freeze_transformer(self, tune_prefix_layers=None):
98 log_str = "Freeze transformer"

Callers 1

__init__Method · 0.45

Calls 2

PromptSpellClass · 0.90
init_method_normalFunction · 0.85

Tested by

no test coverage detected