Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7f579cb351 | ||
|
|
6667726294 |
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user