Compare commits

...
Author SHA1 Message Date
Patrick von Platen d9f6d07c23 delete script 2020-07-17 09:17:46 +00:00
Patrick von Platen afd4037e64 change train script 2020-07-17 07:55:11 +00:00
Patrick von Platen f0d579eb4f adapt params 2020-07-16 12:50:44 +00:00
Patrick von Platen 9e288305ad fix trainer 2020-07-15 17:18:50 +00:00
Patrick von Platen f70d67cfb6 improve script 2020-07-13 18:57:06 +00:00
Patrick von Platen 78166e6dc3 revert wrong changes 2020-07-13 18:34:51 +00:00
Patrick von Platen 524bb46224 train bert encoder decoder 2020-07-13 18:31:27 +00:00
Patrick von Platen e4437a462b first commit 2020-07-10 15:17:49 +00:00
2 changed files with 28 additions and 8 deletions
+22 -8
View File
@@ -820,15 +820,29 @@ class Trainer:
inputs["mems"] = past
with torch.no_grad():
outputs = model(**inputs)
if has_labels:
step_eval_loss, logits = outputs[:2]
eval_losses += [step_eval_loss.mean().item()]
else:
logits = outputs[0]
if self.args.past_index >= 0:
past = outputs[self.args.past_index if has_labels else self.args.past_index - 1]
if self.args.predict_from_generate:
max_length = model.config.max_length
logits_out = model.generate(inputs["input_ids"], attention_mask=inputs["attention_mask"])
# in case the batch is shorter then max length, the output should be padded
logits = model.config.eos_token_id * torch.ones(
(logits_out.shape[0], max_length), dtype=logits_out.dtype, device=logits_out.device
)
logits[:, : logits_out.shape[-1]] = logits_out
if has_labels:
outputs = model(**inputs)
step_eval_loss = outputs[0]
eval_losses += [step_eval_loss.mean().item()]
else:
outputs = model(**inputs)
if has_labels:
step_eval_loss, logits = outputs[:2]
eval_losses += [step_eval_loss.mean().item()]
else:
logits = outputs[0]
if self.args.past_index >= 0:
past = outputs[self.args.past_index if has_labels else self.args.past_index - 1]
if not prediction_loss_only:
if preds is None:
preds = logits.detach()
+6
View File
@@ -157,6 +157,12 @@ class TrainingArguments:
default=1,
metadata={"help": "Number of updates steps to accumulate before performing a backward/update pass."},
)
predict_from_generate: bool = field(
default=False,
metadata={
"help": "Use generate function to predict logits. This is usually the case for summarization or translation."
},
)
learning_rate: float = field(default=5e-5, metadata={"help": "The initial learning rate for Adam."})
weight_decay: float = field(default=0.0, metadata={"help": "Weight decay if we apply some."})