Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d9f6d07c23 | ||
|
|
afd4037e64 | ||
|
|
f0d579eb4f | ||
|
|
9e288305ad | ||
|
|
f70d67cfb6 | ||
|
|
78166e6dc3 | ||
|
|
524bb46224 | ||
|
|
e4437a462b |
@@ -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()
|
||||
|
||||
@@ -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."})
|
||||
|
||||
Reference in New Issue
Block a user