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

Function load_tf_weights_in_bert

model/modeling_bert.py:90–148  ·  view source on GitHub ↗

Load tf checkpoints in a pytorch model

(model, tf_checkpoint_path)

Source from the content-addressed store, hash-verified

88
89
90def load_tf_weights_in_bert(model, tf_checkpoint_path):
91 """ Load tf checkpoints in a pytorch model
92 """
93 try:
94 import re
95 import numpy as np
96 import tensorflow as tf
97 except ImportError:
98 print("Loading a TensorFlow models in PyTorch, requires TensorFlow to be installed. Please see "
99 "https://www.tensorflow.org/install/ for installation instructions.")
100 raise
101 tf_path = os.path.abspath(tf_checkpoint_path)
102 print("Converting TensorFlow checkpoint from {}".format(tf_path))
103 # Load weights from TF model
104 init_vars = tf.train.list_variables(tf_path)
105 names = []
106 arrays = []
107 for name, shape in init_vars:
108 print("Loading TF weight {} with shape {}".format(name, shape))
109 array = tf.train.load_variable(tf_path, name)
110 names.append(name)
111 arrays.append(array)
112
113 for name, array in zip(names, arrays):
114 name = name.split('/')
115 # adam_v and adam_m are variables used in AdamWeightDecayOptimizer to calculated m and v
116 # which are not required for using pretrained model
117 if any(n in ["adam_v", "adam_m"] for n in name):
118 print("Skipping {}".format("/".join(name)))
119 continue
120 pointer = model
121 for m_name in name:
122 if re.fullmatch(r'[A-Za-z]+_\d+', m_name):
123 l = re.split(r'_(\d+)', m_name)
124 else:
125 l = [m_name]
126 if l[0] == 'kernel' or l[0] == 'gamma':
127 pointer = getattr(pointer, 'weight')
128 elif l[0] == 'output_bias' or l[0] == 'beta':
129 pointer = getattr(pointer, 'bias')
130 elif l[0] == 'output_weights':
131 pointer = getattr(pointer, 'weight')
132 else:
133 pointer = getattr(pointer, l[0])
134 if len(l) >= 2:
135 num = int(l[1])
136 pointer = pointer[num]
137 if m_name[-11:] == '_embeddings':
138 pointer = getattr(pointer, 'weight')
139 elif m_name == 'kernel':
140 array = np.transpose(array)
141 try:
142 assert pointer.shape == array.shape
143 except AssertionError as e:
144 e.args += (pointer.shape, array.shape)
145 raise
146 print("Initialize PyTorch weight {}".format(name))
147 pointer.data = torch.from_numpy(array)

Callers

nothing calls this directly

Calls 1

appendMethod · 0.80

Tested by

no test coverage detected