Load tf checkpoints in a pytorch model
(model, tf_checkpoint_path)
| 88 | |
| 89 | |
| 90 | def 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) |