Compare commits

...
5 changed files with 28 additions and 20 deletions
+4 -1
View File
@@ -424,7 +424,10 @@ def main():
eval_dataset = Subset(eval_dataset, list(range(min(args.data_subset, len(eval_dataset)))))
eval_sampler = SequentialSampler(eval_dataset) if args.local_rank == -1 else DistributedSampler(eval_dataset)
eval_dataloader = DataLoader(
eval_dataset, sampler=eval_sampler, batch_size=args.batch_size, collate_fn=DefaultDataCollator().collate_batch
eval_dataset,
sampler=eval_sampler,
batch_size=args.batch_size,
collate_fn=DefaultDataCollator(args.device).collate_batch,
)
# Compute head entropy and importance score
@@ -222,7 +222,7 @@ def main():
train_dataset = get_dataset(data_args, tokenizer=tokenizer) if training_args.do_train else None
eval_dataset = get_dataset(data_args, tokenizer=tokenizer, evaluate=True) if training_args.do_eval else None
data_collator = DataCollatorForLanguageModeling(
tokenizer=tokenizer, mlm=data_args.mlm, mlm_probability=data_args.mlm_probability
tokenizer=tokenizer, mlm=data_args.mlm, mlm_probability=data_args.mlm_probability, device=training_args.device
)
# Initialize our Trainer
+18 -7
View File
@@ -1,6 +1,6 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import Any, Dict, List, NewType, Tuple
from typing import Any, Dict, List, NewType, Tuple, Union
import torch
from torch.nn.utils.rnn import pad_sequence
@@ -8,12 +8,20 @@ from torch.nn.utils.rnn import pad_sequence
from ..tokenization_utils import PreTrainedTokenizer
@dataclass
class DataCollator(ABC):
"""
A `DataCollator` is responsible for batching
and pre-processing samples of data as requested by the training loop.
"""
"""
PyTorch device (torch.device) or str referring to where the generate tensors are placed
- "cpu": Tensors are allocated on the CPU
- "cuda:X": Tensors are allocated on the GPU "X"
"""
device: Union[torch.device, str]
@abstractmethod
def collate_batch(self) -> Dict[str, torch.Tensor]:
"""
@@ -49,20 +57,23 @@ class DefaultDataCollator(DataCollator):
# on the whole batch.
first = features[0]
# Cache to avoid class lookup
device = self.device
# Special handling for labels.
# Ensure that tensor is created with the correct type
# (it should be automatically the case, but let's make sure of it.)
if hasattr(first, "label") and first.label is not None:
if type(first.label) is int:
labels = torch.tensor([f.label for f in features], dtype=torch.long)
labels = torch.tensor([f.label for f in features], device=device, dtype=torch.long)
else:
labels = torch.tensor([f.label for f in features], dtype=torch.float)
labels = torch.tensor([f.label for f in features], device=device, dtype=torch.float)
batch = {"labels": labels}
elif hasattr(first, "label_ids") and first.label_ids is not None:
if type(first.label_ids[0]) is int:
labels = torch.tensor([f.label_ids for f in features], dtype=torch.long)
labels = torch.tensor([f.label_ids for f in features], device=device, dtype=torch.long)
else:
labels = torch.tensor([f.label_ids for f in features], dtype=torch.float)
labels = torch.tensor([f.label_ids for f in features], device=device, dtype=torch.float)
batch = {"labels": labels}
else:
batch = {}
@@ -71,7 +82,7 @@ class DefaultDataCollator(DataCollator):
# Again, we will use the first element to figure out which key/values are not None for this model.
for k, v in vars(first).items():
if k not in ("label", "label_ids") and v is not None and not isinstance(v, str):
batch[k] = torch.tensor([getattr(f, k) for f in features], dtype=torch.long)
batch[k] = torch.tensor([getattr(f, k) for f in features], device=device, dtype=torch.long)
return batch
@@ -99,7 +110,7 @@ class DataCollatorForLanguageModeling(DataCollator):
length_of_first = examples[0].size(0)
are_tensors_same_length = all(x.size(0) == length_of_first for x in examples)
if are_tensors_same_length:
return torch.stack(examples, dim=0)
return torch.stack(examples, dim=0).to(self.device)
else:
if self.tokenizer._pad_token is None:
raise ValueError(
+1 -7
View File
@@ -193,7 +193,7 @@ class Trainer:
if data_collator is not None:
self.data_collator = data_collator
else:
self.data_collator = DefaultDataCollator()
self.data_collator = DefaultDataCollator(args.device)
self.train_dataset = train_dataset
self.eval_dataset = eval_dataset
self.compute_metrics = compute_metrics
@@ -565,9 +565,6 @@ class Trainer:
self, model: nn.Module, inputs: Dict[str, torch.Tensor], optimizer: torch.optim.Optimizer
) -> float:
model.train()
for k, v in inputs.items():
inputs[k] = v.to(self.args.device)
outputs = model(**inputs)
loss = outputs[0] # model outputs are always tuple in transformers (see doc)
@@ -749,9 +746,6 @@ class Trainer:
for inputs in tqdm(dataloader, desc=description):
has_labels = any(inputs.get(k) is not None for k in ["labels", "lm_labels", "masked_lm_labels"])
for k, v in inputs.items():
inputs[k] = v.to(self.args.device)
with torch.no_grad():
outputs = model(**inputs)
if has_labels:
+4 -4
View File
@@ -31,7 +31,7 @@ class DataCollatorIntegrationTest(unittest.TestCase):
task_name="mrpc", data_dir="./tests/fixtures/tests_samples/MRPC", overwrite_cache=True
)
dataset = GlueDataset(data_args, tokenizer=tokenizer, mode="dev")
data_collator = DefaultDataCollator()
data_collator = DefaultDataCollator(device="cpu")
batch = data_collator.collate_batch(dataset.features)
self.assertEqual(batch["labels"].dtype, torch.long)
@@ -42,13 +42,13 @@ class DataCollatorIntegrationTest(unittest.TestCase):
task_name="sts-b", data_dir="./tests/fixtures/tests_samples/STS-B", overwrite_cache=True
)
dataset = GlueDataset(data_args, tokenizer=tokenizer, mode="dev")
data_collator = DefaultDataCollator()
data_collator = DefaultDataCollator(device="cpu")
batch = data_collator.collate_batch(dataset.features)
self.assertEqual(batch["labels"].dtype, torch.float)
def test_lm_tokenizer_without_padding(self):
tokenizer = AutoTokenizer.from_pretrained("gpt2")
data_collator = DataCollatorForLanguageModeling(tokenizer, mlm=False)
data_collator = DataCollatorForLanguageModeling(device="cpu", tokenizer=tokenizer, mlm=False)
# ^ causal lm
dataset = LineByLineTextDataset(tokenizer, file_path=PATH_SAMPLE_TEXT, block_size=512)
@@ -66,7 +66,7 @@ class DataCollatorIntegrationTest(unittest.TestCase):
def test_lm_tokenizer_with_padding(self):
tokenizer = AutoTokenizer.from_pretrained("distilroberta-base")
data_collator = DataCollatorForLanguageModeling(tokenizer)
data_collator = DataCollatorForLanguageModeling(device="cpu", tokenizer=tokenizer)
# ^ masked lm
dataset = LineByLineTextDataset(tokenizer, file_path=PATH_SAMPLE_TEXT, block_size=512)