From 4781afd045b4722e7f28347f1c4f42a56a4550e8 Mon Sep 17 00:00:00 2001 From: Sylvain Gugger <35901082+sgugger@users.noreply.github.com> Date: Mon, 20 Jul 2020 19:47:06 -0400 Subject: [PATCH] Clarify arg class (#5916) --- src/transformers/trainer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/transformers/trainer.py b/src/transformers/trainer.py index 566dd54a..543f6fca 100755 --- a/src/transformers/trainer.py +++ b/src/transformers/trainer.py @@ -144,9 +144,9 @@ class Trainer: data_collator (:obj:`DataCollator`, `optional`, defaults to :func:`~transformers.default_data_collator`): The function to use to from a batch from a list of elements of :obj:`train_dataset` or :obj:`eval_dataset`. - train_dataset (:obj:`Dataset`, `optional`): + train_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`): The dataset to use for training. - eval_dataset (:obj:`Dataset`, `optional`): + eval_dataset (:obj:`torch.utils.data.dataset.Dataset`, `optional`): The dataset to use for evaluation. compute_metrics (:obj:`Callable[[EvalPrediction], Dict]`, `optional`): The function that will be used to compute metrics at evaluation. Must take a