MCPcopy Create free account
hub / github.com/PAIR-code/lit / train

Method train

lit_nlp/examples/glue/models.py:185–229  ·  view source on GitHub ↗

Run fine-tuning.

(
      self,
      train_inputs: list[JsonDict],
      validation_inputs: list[JsonDict],
      learning_rate=2e-5,
      batch_size=32,
      num_epochs=3,
      keras_callbacks=None,
  )

Source from the content-addressed store, hash-verified

183 return tf.data.Dataset.from_tensor_slices((dict(encoded_input), labels))
184
185 def train(
186 self,
187 train_inputs: list[JsonDict],
188 validation_inputs: list[JsonDict],
189 learning_rate=2e-5,
190 batch_size=32,
191 num_epochs=3,
192 keras_callbacks=None,
193 ):
194 """Run fine-tuning."""
195 train_dataset = (
196 self._make_dataset(train_inputs)
197 .shuffle(128)
198 .batch(batch_size)
199 .repeat(-1)
200 )
201 # Use larger batch for validation since inference is about 1/2 memory usage
202 # of backprop.
203 eval_batch_size = 2 * batch_size
204 validation_dataset = self._make_dataset(validation_inputs).batch(
205 eval_batch_size
206 )
207
208 # Prepare model for training.
209 opt = keras.optimizers.Adam(learning_rate=learning_rate, epsilon=1e-08)
210 if self.is_regression:
211 loss = keras.losses.MeanSquaredError()
212 metric = keras.metrics.RootMeanSquaredError("rmse")
213 else:
214 loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True)
215 metric = keras.metrics.SparseCategoricalAccuracy("accuracy")
216 self.model.compile(optimizer=opt, loss=loss, metrics=[metric])
217
218 steps_per_epoch = len(train_inputs) // batch_size
219 validation_steps = len(validation_inputs) // eval_batch_size
220 history = self.model.fit(
221 train_dataset,
222 epochs=num_epochs,
223 steps_per_epoch=steps_per_epoch,
224 validation_data=validation_dataset,
225 validation_steps=validation_steps,
226 callbacks=keras_callbacks,
227 verbose=2,
228 )
229 return history
230
231 def save(self, path: str):
232 """Save model weights and tokenizer info.

Callers 1

train_and_saveFunction · 0.80

Calls 2

_make_datasetMethod · 0.95
shuffleMethod · 0.80

Tested by

no test coverage detected