Compare commits
21
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e50654c03c | ||
|
|
218ae1ed18 | ||
|
|
50fd3479f3 | ||
|
|
ce1e6c3d74 | ||
|
|
343157dd96 | ||
|
|
0c75018138 | ||
|
|
1477c8aa89 | ||
|
|
1b1b94ed8f | ||
|
|
edb64e8ac3 | ||
|
|
f49448aa15 | ||
|
|
b5b9633e2c | ||
|
|
f842f7c764 | ||
|
|
ea3e2cebd8 | ||
|
|
a6f5989bd9 | ||
|
|
3488e5a8b8 | ||
|
|
f1c91f8fad | ||
|
|
75af46eada | ||
|
|
ac396b2787 | ||
|
|
75d900e7c4 | ||
|
|
91dc392e57 | ||
|
|
30b2dbbba6 |
@@ -0,0 +1,756 @@
|
||||
# coding=utf-8
|
||||
# Copyright 2018 The Google AI Language Team Authors and The HuggingFace Inc. team.
|
||||
# Copyright (c) 2018, NVIDIA CORPORATION. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""
|
||||
Pre-training language models using the ELECTRA method.
|
||||
"""
|
||||
|
||||
|
||||
import logging
|
||||
import math
|
||||
import multiprocessing
|
||||
import os
|
||||
import tarfile
|
||||
from dataclasses import dataclass, field
|
||||
from itertools import chain
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.utils.data import IterableDataset
|
||||
from tqdm import tqdm
|
||||
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
AutoTokenizer,
|
||||
DataCollatorForLanguageModeling,
|
||||
ElectraForMaskedLM,
|
||||
ElectraForPreTraining,
|
||||
EvalPrediction,
|
||||
HfArgumentParser,
|
||||
PreTrainedTokenizer,
|
||||
TextDataset,
|
||||
Trainer,
|
||||
set_seed,
|
||||
)
|
||||
from transformers.modeling_utils import PreTrainedModel
|
||||
from transformers.training_args import TrainingArguments
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelArguments:
|
||||
"""
|
||||
Arguments pertaining to which model/config/tokenizer we are going to fine-tune, or train from scratch.
|
||||
"""
|
||||
|
||||
discriminator_name_or_path: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "The discriminator checkpoint for weights initialization. Leave None if you want to train a model "
|
||||
"from scratch."
|
||||
},
|
||||
)
|
||||
|
||||
discriminator_config_name: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={"help": "Pretrained config name or path if not the same as the discriminator model_name"},
|
||||
)
|
||||
|
||||
generator_name_or_path: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={
|
||||
"help": "The generator checkpoint for weights initialization. Leave None if you want to train a model "
|
||||
"from scratch."
|
||||
},
|
||||
)
|
||||
|
||||
generator_config_name: Optional[str] = field(
|
||||
default=None, metadata={"help": "Pretrained config name or path if not the same as the generator model_name"}
|
||||
)
|
||||
|
||||
tokenizer_name: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={"help": "Pretrained tokenizer name or path if not the same as discriminator model_name"},
|
||||
)
|
||||
|
||||
cache_dir: Optional[str] = field(
|
||||
default=None, metadata={"help": "Where do you want to store the pretrained models downloaded from s3"}
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataTrainingArguments:
|
||||
"""
|
||||
Arguments pertaining to what data we are going to input our model for training and eval.
|
||||
"""
|
||||
|
||||
train_data_file: Optional[str] = field(
|
||||
default=None, metadata={"help": "The input training data file (a text file)."}
|
||||
)
|
||||
|
||||
eval_data_file: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={"help": "An optional input evaluation data file to evaluate the perplexity on (a text file)."},
|
||||
)
|
||||
|
||||
open_web_text_directory: Optional[str] = field(
|
||||
default=None,
|
||||
metadata={"help": "The directory containing files that will be used for training and evaluation."},
|
||||
)
|
||||
|
||||
block_size: int = field(
|
||||
default=-1,
|
||||
metadata={
|
||||
"help": "Optional input sequence length after tokenization."
|
||||
"The training dataset will be truncated in block of this size for training."
|
||||
"Default to the model max input length for single sentence inputs (take into account special tokens)."
|
||||
},
|
||||
)
|
||||
|
||||
overwrite_cache: bool = field(
|
||||
default=False, metadata={"help": "Overwrite the cached training and evaluation sets"}
|
||||
)
|
||||
|
||||
num_dataset_building_processes: int = field(
|
||||
default=1, metadata={"help": "The number of workers that will be used to build the dataset."}
|
||||
)
|
||||
|
||||
num_tensors_per_file: int = field(
|
||||
default=2048,
|
||||
metadata={
|
||||
"help": "The number of tensors that will be stored in each file after tokenization."
|
||||
"The smaller the amount, the smaller the filesize, but the larger the amount"
|
||||
"of files that will be created."
|
||||
},
|
||||
)
|
||||
|
||||
mask_probability: float = field(
|
||||
default=0.15, metadata={"help": "Percentage of the input that will be masked or replaced."}
|
||||
)
|
||||
|
||||
max_predictions_per_sequence: int = field(
|
||||
default=-1, metadata={"help": "Maximum tokens that will be masked in a sequence."},
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ElectraTrainingArguments(TrainingArguments):
|
||||
max_steps: int = field(
|
||||
default=1_000_000,
|
||||
metadata={"help": "If > 0: set total number of training steps to perform. Override num_train_epochs."},
|
||||
)
|
||||
|
||||
max_eval_steps: int = field(
|
||||
default=100, metadata={"help": "If > 0: set total number of eval steps to perform."},
|
||||
)
|
||||
warmup_steps: int = field(default=10_000, metadata={"help": "Linear warmup over warmup_steps."})
|
||||
|
||||
weight_decay: float = field(default=0.1, metadata={"help": "Weight decay if we apply some."})
|
||||
|
||||
generator_weight: float = field(default=1.0, metadata={"help": "Weight coefficient for the generator loss"})
|
||||
|
||||
discriminator_weight: float = field(
|
||||
default=50.0, metadata={"help": "Weight coefficient for the discriminator loss"}
|
||||
)
|
||||
|
||||
|
||||
def get_dataset(
|
||||
data_args: DataTrainingArguments,
|
||||
training_args: TrainingArguments,
|
||||
model_args: ModelArguments,
|
||||
tokenizer: Union[PreTrainedTokenizer, str],
|
||||
evaluate=False,
|
||||
local_rank=-1,
|
||||
):
|
||||
if data_args.open_web_text_directory is not None:
|
||||
# Whether to overwrite the cache. We don't want to overwrite the cache when evaluating if we're both training
|
||||
# and evaluating, as the same dataset is used. We don't need to tokenize twice for training
|
||||
# and evaluation.
|
||||
should_overwrite_cache = False
|
||||
|
||||
# If argument is specified and training, then respect the argument
|
||||
if data_args.overwrite_cache and not evaluate:
|
||||
should_overwrite_cache = True
|
||||
if data_args.overwrite_cache and training_args.do_eval and not training_args.do_train:
|
||||
should_overwrite_cache = True
|
||||
|
||||
return OpenWebTextDataset(data_args, model_args, overwrite_cache=should_overwrite_cache)
|
||||
else:
|
||||
file_path = data_args.eval_data_file if evaluate else data_args.train_data_file
|
||||
return TextDataset(
|
||||
tokenizer=tokenizer, file_path=file_path, block_size=data_args.block_size, local_rank=local_rank,
|
||||
)
|
||||
|
||||
|
||||
class OpenWebTextDataset(IterableDataset):
|
||||
def __init__(self, data_args, model_args, overwrite_cache=False):
|
||||
self.tokenizer_cache = model_args.cache_dir
|
||||
self.directory = Path(data_args.open_web_text_directory)
|
||||
self.archives = os.listdir(self.directory)
|
||||
self.tokenizer_identifier = model_args.tokenizer_name
|
||||
self.num_tensors_per_file = data_args.num_tensors_per_file
|
||||
self.feature_directory = (
|
||||
self.directory
|
||||
/ f"features_{self.tokenizer_identifier.replace('/', '_')}_{data_args.block_size if data_args.block_size is not None else 'no-max-seq'}_{self.num_tensors_per_file}"
|
||||
)
|
||||
self.block_size = data_args.block_size
|
||||
|
||||
# The dataset was already processed
|
||||
if os.path.exists(self.feature_directory) and not overwrite_cache:
|
||||
logger.info(
|
||||
f"Re-using cache from {self.feature_directory}. Warning: we have no way of detecting an "
|
||||
f"incomplete cache. If the tokenization was started but not finished, please use the "
|
||||
f"`--ignore_cache=True` flag."
|
||||
)
|
||||
self.feature_set_paths = [
|
||||
self.feature_directory / feature_set_path for feature_set_path in os.listdir(self.feature_directory)
|
||||
]
|
||||
return
|
||||
|
||||
logger.info(f"Writing features at {self.feature_directory}")
|
||||
os.makedirs(self.feature_directory, exist_ok=overwrite_cache)
|
||||
|
||||
n_archives_per_job = math.ceil(len(self.archives) / data_args.num_dataset_building_processes)
|
||||
self.job_archives = [
|
||||
self.archives[i * n_archives_per_job : (i + 1) * n_archives_per_job]
|
||||
for i in range(data_args.num_dataset_building_processes)
|
||||
]
|
||||
# Sanity check: make sure we're not leaving any archive behind.
|
||||
assert sum([len(archive) for archive in self.job_archives]) == len(self.archives)
|
||||
|
||||
if data_args.num_dataset_building_processes == 1:
|
||||
self.feature_set_paths = self._extract_open_web_text()
|
||||
else:
|
||||
pool = multiprocessing.Pool(processes=data_args.num_dataset_building_processes)
|
||||
self.feature_set_paths = pool.map(
|
||||
self._extract_open_web_text, range(data_args.num_dataset_building_processes)
|
||||
)
|
||||
self.feature_set_paths = [file_path for feature_set in self.feature_set_paths for file_path in feature_set]
|
||||
|
||||
def _extract_open_web_text(self, job_id=0):
|
||||
"""
|
||||
OpenWebText is saved under the following format:
|
||||
|
||||
openwebtext.zip
|
||||
|-> archive_xxx.zip
|
||||
|-> file_xxx.txt
|
||||
|-> file_xxz.txt
|
||||
...
|
||||
|-> archive_xxz.zip
|
||||
|-> file_xxy.txt
|
||||
...
|
||||
...
|
||||
"""
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
self.tokenizer_identifier, use_fast=True, cache_dir=self.tokenizer_cache
|
||||
)
|
||||
|
||||
# Create openwebtext/tmp directory to store temporary files
|
||||
temporary_directory = self.directory / "tmp" / f"job_{job_id}"
|
||||
feature_index = 0
|
||||
feature_set_paths = []
|
||||
|
||||
os.makedirs(temporary_directory, exist_ok=True)
|
||||
|
||||
# Extract archives and tokenize in directory
|
||||
progress_bar = tqdm(
|
||||
self.job_archives[job_id], desc="Extracting archives", total=len(self.job_archives[0]), disable=job_id != 0
|
||||
)
|
||||
|
||||
features = []
|
||||
for archive in progress_bar:
|
||||
if os.path.isdir(self.directory / archive):
|
||||
logger.info("Ignoring rogue directory.")
|
||||
continue
|
||||
with tarfile.open(self.directory / archive) as t:
|
||||
extracted_archive = temporary_directory / f"{archive}-extracted"
|
||||
t.extractall(extracted_archive)
|
||||
|
||||
files = os.listdir(extracted_archive)
|
||||
for file in files:
|
||||
file_path = extracted_archive / file
|
||||
|
||||
with open(file_path, "r") as f:
|
||||
text = f.read()
|
||||
block_size = tokenizer.model_max_length if self.block_size is None else self.block_size
|
||||
encoding = tokenizer.encode_plus(text, return_overflowing_tokens=True, max_length=block_size)
|
||||
|
||||
features.append(torch.tensor(encoding["input_ids"]))
|
||||
|
||||
for overflowing_encoding in encoding.encodings[0].overflowing:
|
||||
features.append(torch.tensor(overflowing_encoding.ids))
|
||||
|
||||
while len(features) > self.num_tensors_per_file:
|
||||
feature_set_path = self.feature_directory / f"feature_set_{job_id}_{feature_index}.pt"
|
||||
torch.save(features[: self.num_tensors_per_file], feature_set_path)
|
||||
features = features[self.num_tensors_per_file :]
|
||||
feature_index += 1
|
||||
feature_set_paths.append(feature_set_path)
|
||||
|
||||
if len(features) > 0:
|
||||
feature_set_path = self.feature_directory / f"feature_set_{job_id}_{feature_index}.pt"
|
||||
torch.save(features, feature_set_path)
|
||||
feature_set_paths.append(feature_set_path)
|
||||
|
||||
return feature_set_paths
|
||||
|
||||
@staticmethod
|
||||
def parse_file(file_index):
|
||||
try:
|
||||
features = torch.load(file_index)
|
||||
yield from features
|
||||
except RuntimeError:
|
||||
raise RuntimeError(f"Corrupted file {file_index}")
|
||||
|
||||
def __len__(self):
|
||||
return len(self.feature_set_paths) * self.num_tensors_per_file
|
||||
|
||||
def __iter__(self):
|
||||
return chain.from_iterable(map(self.parse_file, self.feature_set_paths))
|
||||
|
||||
|
||||
class CombinedModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
discriminator: PreTrainedModel,
|
||||
generator: PreTrainedModel,
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
training_args: ElectraTrainingArguments,
|
||||
data_args: DataTrainingArguments,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.discriminator = discriminator
|
||||
self.generator = generator
|
||||
|
||||
# Embeddings are shared
|
||||
self.discriminator.set_input_embeddings(self.generator.get_input_embeddings())
|
||||
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
self.discriminator_weight = training_args.discriminator_weight
|
||||
self.generator_weight = training_args.generator_weight
|
||||
self.mask_probability = data_args.mask_probability
|
||||
self.max_predictions_per_sequence = data_args.max_predictions_per_sequence
|
||||
|
||||
# Original implementation has a default of :(mask_probability + 0.005) * max_sequence_length
|
||||
if self.max_predictions_per_sequence == -1:
|
||||
self.max_predictions_per_sequence = (self.mask_probability + 0.005) * data_args.block_size
|
||||
|
||||
class Config:
|
||||
xla_device: bool = False
|
||||
|
||||
self.config = Config()
|
||||
|
||||
def mask_inputs(
|
||||
self, input_ids: torch.Tensor, tokens_to_ignore, proposal_distribution=1.0,
|
||||
):
|
||||
input_ids = input_ids.clone()
|
||||
inputs_which_can_be_masked = torch.ones_like(input_ids)
|
||||
for token in tokens_to_ignore:
|
||||
inputs_which_can_be_masked -= torch.eq(input_ids, token).long()
|
||||
|
||||
total_number_of_tokens = input_ids.shape[-1]
|
||||
|
||||
# Identify the number of tokens to be masked, which should be: 1 < num < max_predictions per seq.
|
||||
# It is set to be: n_tokens * mask_probability, but is truncated if it goes beyond bounds.
|
||||
number_of_tokens_to_be_masked = torch.max(
|
||||
torch.tensor(1),
|
||||
torch.min(
|
||||
torch.tensor(self.max_predictions_per_sequence, dtype=torch.long),
|
||||
torch.tensor(int(total_number_of_tokens * self.mask_probability), dtype=torch.long),
|
||||
),
|
||||
)
|
||||
|
||||
# The probability of each token being masked
|
||||
sample_prob = proposal_distribution * inputs_which_can_be_masked
|
||||
sample_prob /= torch.sum(sample_prob)
|
||||
|
||||
# Sample from the probabilities
|
||||
masked_lm_positions = sample_prob.multinomial(number_of_tokens_to_be_masked)
|
||||
|
||||
# Gather the IDs from the positions
|
||||
masked_lm_ids = input_ids.gather(-1, masked_lm_positions)
|
||||
|
||||
return masked_lm_ids, masked_lm_positions
|
||||
|
||||
@staticmethod
|
||||
def gather_positions(sequence, positions):
|
||||
batch_size, sequence_length, dimension = sequence.shape
|
||||
position_shift = (sequence_length * torch.arange(batch_size, device=sequence.device)).unsqueeze(-1)
|
||||
flat_positions = torch.reshape(positions + position_shift, [-1]).long()
|
||||
flat_sequence = torch.reshape(sequence, [batch_size * sequence_length, dimension])
|
||||
gathered = flat_sequence.index_select(0, flat_positions)
|
||||
return torch.reshape(gathered, [batch_size, -1, dimension])
|
||||
|
||||
def forward(
|
||||
self, input_ids=None, attention_mask=None, token_type_ids=None, position_ids=None, head_mask=None, labels=None,
|
||||
):
|
||||
# get the masked positions as well as their original values
|
||||
masked_input_ids, masked_lm_positions = self.mask_inputs(
|
||||
input_ids, [self.tokenizer.cls_token_id, self.tokenizer.sep_token_id, self.tokenizer.mask_token_id],
|
||||
)
|
||||
|
||||
# only masked values should be counted in the loss; build a tensor containing the true values and -100 otherwise
|
||||
masked_lm_labels = torch.full_like(input_ids, -100)
|
||||
masked_lm_labels.scatter_(-1, masked_lm_positions, masked_input_ids)
|
||||
|
||||
# Create a tensor filled with masks
|
||||
masked_tokens = torch.full_like(masked_input_ids, self.tokenizer.mask_token_id)
|
||||
masked_lm_inputs = input_ids.clone()
|
||||
|
||||
# Of the evaluated tokens, 15% of those will keep their original tokens
|
||||
replace_with_mask_positions = masked_lm_positions * (
|
||||
torch.rand(masked_lm_positions.shape, device=masked_lm_positions.device) < (1 - self.mask_probability)
|
||||
)
|
||||
|
||||
# Scatter the masks at the masked positions
|
||||
masked_lm_inputs.scatter_(-1, replace_with_mask_positions, masked_tokens)
|
||||
masked_lm_inputs[..., 0] = 101
|
||||
|
||||
generator_loss, generator_output = self.generator(
|
||||
masked_lm_inputs,
|
||||
attention_mask,
|
||||
token_type_ids,
|
||||
position_ids,
|
||||
head_mask,
|
||||
position_ids,
|
||||
masked_lm_labels=masked_lm_labels,
|
||||
)[:2]
|
||||
|
||||
# get the generator's predicted value on each masked position
|
||||
fake_logits = self.gather_positions(generator_output, masked_lm_positions)
|
||||
fake_softmaxed = torch.softmax(fake_logits, dim=-1)
|
||||
fake_sampled = fake_softmaxed\
|
||||
.view(fake_logits.shape[0] * fake_logits.shape[1], fake_logits.shape[2])\
|
||||
.multinomial(1)\
|
||||
.view(fake_logits.shape[:-1])
|
||||
|
||||
# create a tensor containing the predicted tokens
|
||||
fake_tokens = input_ids.scatter(-1, masked_lm_positions, fake_sampled)
|
||||
fake_tokens[:, 0] = input_ids[:, 0]
|
||||
discriminator_labels = (labels != fake_tokens).int()
|
||||
|
||||
discriminator_loss, discriminator_output = self.discriminator(
|
||||
fake_tokens,
|
||||
attention_mask,
|
||||
token_type_ids,
|
||||
position_ids,
|
||||
head_mask,
|
||||
position_ids,
|
||||
labels=discriminator_labels,
|
||||
)[:2]
|
||||
|
||||
discriminator_predictions = torch.round((torch.sign(discriminator_output) + 1.0) * 0.5)
|
||||
|
||||
total_loss = (self.discriminator_weight * discriminator_loss) + (self.generator_weight * generator_loss)
|
||||
|
||||
return (
|
||||
total_loss,
|
||||
(generator_output, discriminator_output),
|
||||
(masked_input_ids, fake_sampled),
|
||||
(discriminator_labels, discriminator_predictions),
|
||||
)
|
||||
|
||||
def save_pretrained(self, directory):
|
||||
if self.config.xla_device:
|
||||
self.discriminator.config.xla_device = True
|
||||
self.generator.config.xla_device = True
|
||||
else:
|
||||
self.discriminator.config.xla_device = False
|
||||
self.generator.config.xla_device = False
|
||||
|
||||
generator_path = os.path.join(directory, "generator")
|
||||
discriminator_path = os.path.join(directory, "discriminator")
|
||||
|
||||
if not os.path.exists(generator_path):
|
||||
os.makedirs(generator_path)
|
||||
|
||||
if not os.path.exists(discriminator_path):
|
||||
os.makedirs(discriminator_path)
|
||||
|
||||
self.generator.save_pretrained(generator_path)
|
||||
self.discriminator.save_pretrained(discriminator_path)
|
||||
|
||||
|
||||
def main():
|
||||
# See all possible arguments in src/transformers/training_args.py
|
||||
# or by passing the --help flag to this script.
|
||||
# We now keep distinct sets of args, for a cleaner separation of concerns.
|
||||
|
||||
parser = HfArgumentParser((ModelArguments, DataTrainingArguments, ElectraTrainingArguments))
|
||||
model_args, data_args, training_args = parser.parse_args_into_dataclasses()
|
||||
|
||||
if data_args.open_web_text_directory is None and data_args.eval_data_file is None and training_args.do_eval:
|
||||
raise ValueError(
|
||||
"Cannot do evaluation without an evaluation data file. Either supply a file to --eval_data_file "
|
||||
"or remove the --do_eval argument."
|
||||
)
|
||||
|
||||
if (
|
||||
os.path.exists(training_args.output_dir)
|
||||
and os.listdir(training_args.output_dir)
|
||||
and training_args.do_train
|
||||
and not training_args.overwrite_output_dir
|
||||
):
|
||||
raise ValueError(
|
||||
f"Output directory ({training_args.output_dir}) already exists and is not empty. Use --overwrite_output_dir to overcome."
|
||||
)
|
||||
|
||||
# Setup logging
|
||||
logging.basicConfig(
|
||||
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||
datefmt="%m/%d/%Y %H:%M:%S",
|
||||
level=logging.INFO if training_args.local_rank in [-1, 0] else logging.WARN,
|
||||
)
|
||||
logger.warning(
|
||||
"Process rank: %s, device: %s, n_gpu: %s, distributed training: %s, 16-bits training: %s",
|
||||
training_args.local_rank,
|
||||
training_args.device,
|
||||
training_args.n_gpu,
|
||||
bool(training_args.local_rank != -1),
|
||||
training_args.fp16,
|
||||
)
|
||||
logger.info("Training/evaluation parameters %s", training_args)
|
||||
|
||||
# Set seed
|
||||
set_seed(training_args.seed)
|
||||
|
||||
# Load pretrained model and tokenizer
|
||||
#
|
||||
# Distributed training:
|
||||
# The .from_pretrained methods guarantee that only one local process can concurrently
|
||||
# download model & vocab.
|
||||
|
||||
if model_args.discriminator_config_name:
|
||||
discriminator_config = AutoConfig.from_pretrained(
|
||||
model_args.discriminator_config_name, cache_dir=model_args.cache_dir
|
||||
)
|
||||
elif model_args.discriminator_name_or_path:
|
||||
discriminator_config = AutoConfig.from_pretrained(
|
||||
model_args.discriminator_name_or_path, cache_dir=model_args.cache_dir
|
||||
)
|
||||
else:
|
||||
raise ValueError("Either --discriminator_config_name or --discriminator_name_or_path should be specified.")
|
||||
|
||||
if model_args.generator_config_name:
|
||||
generator_config = AutoConfig.from_pretrained(model_args.generator_config_name, cache_dir=model_args.cache_dir)
|
||||
elif model_args.generator_name_or_path:
|
||||
generator_config = AutoConfig.from_pretrained(
|
||||
model_args.generator_name_or_path, cache_dir=model_args.cache_dir
|
||||
)
|
||||
else:
|
||||
raise ValueError("Either --generator_config_name or --generator_name_or_path should be specified.")
|
||||
|
||||
if model_args.tokenizer_name:
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_args.tokenizer_name, cache_dir=model_args.cache_dir)
|
||||
elif model_args.discriminator_name_or_path:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_args.discriminator_name_or_path, cache_dir=model_args.cache_dir
|
||||
)
|
||||
model_args.tokenizer_name = model_args.discriminator_name_or_path
|
||||
else:
|
||||
raise ValueError(
|
||||
"You are instantiating a new tokenizer from scratch. This is not supported, but you can do it from another "
|
||||
"script, save it, and load it from here, using --tokenizer_name"
|
||||
)
|
||||
|
||||
if model_args.discriminator_name_or_path:
|
||||
discriminator = ElectraForPreTraining.from_pretrained(
|
||||
model_args.discriminator_name_or_path,
|
||||
from_tf=bool(".ckpt" in model_args.discriminator_name_or_path),
|
||||
config=discriminator_config,
|
||||
cache_dir=model_args.cache_dir,
|
||||
)
|
||||
else:
|
||||
logger.info("Training new model from scratch")
|
||||
discriminator = ElectraForPreTraining(discriminator_config)
|
||||
|
||||
if model_args.generator_name_or_path:
|
||||
generator = ElectraForMaskedLM.from_pretrained(
|
||||
model_args.generator_name_or_path,
|
||||
from_tf=bool(".ckpt" in model_args.generator_name_or_path),
|
||||
config=generator_config,
|
||||
cache_dir=model_args.cache_dir,
|
||||
)
|
||||
else:
|
||||
logger.info("Training new model from scratch")
|
||||
generator = ElectraForMaskedLM(generator_config)
|
||||
|
||||
discriminator.resize_token_embeddings(len(tokenizer))
|
||||
generator.resize_token_embeddings(len(tokenizer))
|
||||
|
||||
if data_args.block_size <= 0:
|
||||
data_args.block_size = tokenizer.max_len
|
||||
# Our input block size will be the max possible for the model
|
||||
else:
|
||||
data_args.block_size = min(data_args.block_size, tokenizer.max_len)
|
||||
|
||||
# Need to update this to something cleaner
|
||||
try:
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
if xm.is_master_ordinal(local=True):
|
||||
get_dataset(data_args, training_args, model_args, tokenizer=tokenizer, local_rank=training_args.local_rank)
|
||||
|
||||
xm.rendezvous("dataset building")
|
||||
except ImportError:
|
||||
logger.info("Not running on TPU")
|
||||
|
||||
# Get datasets
|
||||
train_dataset = (
|
||||
get_dataset(data_args, training_args, model_args, tokenizer=tokenizer, local_rank=training_args.local_rank)
|
||||
if training_args.do_train
|
||||
else None
|
||||
)
|
||||
eval_dataset = (
|
||||
get_dataset(
|
||||
data_args,
|
||||
training_args,
|
||||
model_args,
|
||||
tokenizer=tokenizer,
|
||||
local_rank=training_args.local_rank,
|
||||
evaluate=True,
|
||||
)
|
||||
if training_args.do_eval
|
||||
else None
|
||||
)
|
||||
|
||||
# Masking is done inside the CombinedModel
|
||||
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False, mlm_probability=0)
|
||||
|
||||
model = CombinedModel(discriminator, generator, tokenizer, training_args, data_args)
|
||||
|
||||
def compute_metrics(evaluation_predictions: EvalPrediction) -> Dict[str, int]:
|
||||
predictions: Dict[str, np.ndarray] = evaluation_predictions.predictions
|
||||
labels: Dict[str, np.ndarray] = evaluation_predictions.label_ids
|
||||
|
||||
generator_labels, generator_predictions = labels["generator"], predictions["generator"]
|
||||
discriminator_labels, discriminator_predictions = labels["discriminator"], predictions["discriminator"]
|
||||
|
||||
true_positives = (
|
||||
np.logical_and(
|
||||
np.equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 1)
|
||||
)
|
||||
.sum()
|
||||
.astype(float)
|
||||
)
|
||||
false_negatives = (
|
||||
np.logical_and(
|
||||
np.not_equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 1)
|
||||
)
|
||||
.sum()
|
||||
.astype(float)
|
||||
)
|
||||
false_positives = (
|
||||
np.logical_and(
|
||||
np.not_equal(discriminator_predictions, discriminator_labels), np.equal(discriminator_labels, 0)
|
||||
)
|
||||
.sum()
|
||||
.astype(float)
|
||||
)
|
||||
|
||||
generator_accuracy = (
|
||||
np.equal(generator_labels, generator_predictions).sum().astype(float) / generator_predictions.size
|
||||
)
|
||||
discriminator_accuracy = (
|
||||
np.equal(discriminator_labels, discriminator_predictions).sum().astype(float) / discriminator_labels.size
|
||||
)
|
||||
discriminator_precision = true_positives / (true_positives + false_positives)
|
||||
discriminator_recall = true_positives / (true_positives + false_negatives)
|
||||
|
||||
return {
|
||||
"generator_accuracy": generator_accuracy,
|
||||
"discriminator_accuracy": discriminator_accuracy,
|
||||
"discriminator_precision": discriminator_precision,
|
||||
"discriminator_recall": discriminator_recall,
|
||||
}
|
||||
|
||||
def manage_evaluation_predictions(model_outputs: Tuple[torch.Tensor]) -> Dict[str, torch.Tensor]:
|
||||
total_loss, models_output, generator_evaluation_values, discriminator_evaluation_values = model_outputs
|
||||
generator_labels, generator_predictions = generator_evaluation_values
|
||||
discriminator_labels, discriminator_predictions = discriminator_evaluation_values
|
||||
|
||||
return {
|
||||
"generator_labels": generator_labels.detach(),
|
||||
"generator_predictions": generator_predictions.detach(),
|
||||
"discriminator_labels": discriminator_labels.detach(),
|
||||
"discriminator_predictions": discriminator_predictions.detach(),
|
||||
}
|
||||
|
||||
def eval_prediction_mapping(dictionary: Dict[str, np.ndarray]) -> EvalPrediction:
|
||||
labels = {"generator": dictionary["generator_labels"], "discriminator": dictionary["discriminator_labels"]}
|
||||
predictions = {
|
||||
"generator": dictionary["generator_predictions"],
|
||||
"discriminator": dictionary["discriminator_predictions"],
|
||||
}
|
||||
|
||||
return EvalPrediction(labels, predictions)
|
||||
|
||||
# Initialize our Trainer
|
||||
trainer = Trainer(
|
||||
model=model,
|
||||
args=training_args,
|
||||
data_collator=data_collator,
|
||||
train_dataset=train_dataset,
|
||||
eval_dataset=eval_dataset,
|
||||
compute_metrics=compute_metrics,
|
||||
manage_evaluation_predictions=manage_evaluation_predictions,
|
||||
eval_prediction_mapping=eval_prediction_mapping,
|
||||
prediction_loss_only=True,
|
||||
)
|
||||
|
||||
# Training
|
||||
if training_args.do_train:
|
||||
model_path = (
|
||||
model_args.discriminator_name_or_path
|
||||
if model_args.discriminator_name_or_path is not None
|
||||
and os.path.isdir(model_args.discriminator_name_or_path)
|
||||
else None
|
||||
)
|
||||
trainer.train(model_path=model_path)
|
||||
trainer.save_model()
|
||||
|
||||
# Evaluation
|
||||
results = {}
|
||||
if training_args.do_eval and training_args.local_rank in [-1, 0]:
|
||||
logger.info("*** Evaluate ***")
|
||||
|
||||
result = trainer.evaluate()
|
||||
|
||||
output_eval_file = os.path.join(training_args.output_dir, "eval_results_lm.txt")
|
||||
with open(output_eval_file, "w") as writer:
|
||||
logger.info("***** Eval results *****")
|
||||
for key in sorted(result.keys()):
|
||||
logger.info(" %s = %s", key, str(result[key]))
|
||||
writer.write("%s = %s\n" % (key, str(result[key])))
|
||||
|
||||
results.update(result)
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def _mp_fn(index):
|
||||
# For xla_spawn (TPUs)
|
||||
main()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -132,7 +132,7 @@ class ElectraDiscriminatorPredictions(nn.Module):
|
||||
def forward(self, discriminator_hidden_states, attention_mask):
|
||||
hidden_states = self.dense(discriminator_hidden_states)
|
||||
hidden_states = get_activation(self.config.hidden_act)(hidden_states)
|
||||
logits = self.dense_prediction(hidden_states).squeeze()
|
||||
logits = self.dense_prediction(hidden_states).squeeze(-1)
|
||||
|
||||
return logits
|
||||
|
||||
|
||||
+92
-38
@@ -14,7 +14,7 @@ import torch
|
||||
from packaging import version
|
||||
from torch import nn
|
||||
from torch.utils.data.dataloader import DataLoader
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from torch.utils.data.dataset import Dataset, IterableDataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from torch.utils.data.sampler import RandomSampler, Sampler, SequentialSampler
|
||||
from tqdm.auto import tqdm, trange
|
||||
@@ -71,6 +71,8 @@ try:
|
||||
_has_wandb = False if os.getenv("WANDB_DISABLED") else True
|
||||
except ImportError:
|
||||
_has_wandb = False
|
||||
except AttributeError:
|
||||
_has_wandb = False
|
||||
|
||||
|
||||
def is_wandb_available():
|
||||
@@ -162,6 +164,8 @@ class Trainer:
|
||||
train_dataset: Optional[Dataset]
|
||||
eval_dataset: Optional[Dataset]
|
||||
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None
|
||||
manage_evaluation_predictions: Optional[Callable[[Tuple[torch.Tensor]], Dict[str, torch.Tensor]]] = None
|
||||
eval_prediction_mapping: Optional[Callable[[Dict[str, np.ndarray]], EvalPrediction]]
|
||||
prediction_loss_only: bool
|
||||
tb_writer: Optional["SummaryWriter"] = None
|
||||
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None
|
||||
@@ -176,6 +180,8 @@ class Trainer:
|
||||
train_dataset: Optional[Dataset] = None,
|
||||
eval_dataset: Optional[Dataset] = None,
|
||||
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None,
|
||||
manage_evaluation_predictions: Optional[Callable[[Tuple[torch.Tensor]], Dict[str, torch.Tensor]]] = None,
|
||||
eval_prediction_mapping: Optional[Callable[[Dict[str, np.ndarray]], EvalPrediction]] = None,
|
||||
prediction_loss_only=False,
|
||||
tb_writer: Optional["SummaryWriter"] = None,
|
||||
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = None,
|
||||
@@ -197,6 +203,8 @@ class Trainer:
|
||||
self.train_dataset = train_dataset
|
||||
self.eval_dataset = eval_dataset
|
||||
self.compute_metrics = compute_metrics
|
||||
self.manage_evaluation_predictions = manage_evaluation_predictions
|
||||
self.eval_prediction_mapping = eval_prediction_mapping
|
||||
self.prediction_loss_only = prediction_loss_only
|
||||
self.optimizers = optimizers
|
||||
if tb_writer is not None:
|
||||
@@ -238,7 +246,7 @@ class Trainer:
|
||||
data_loader = DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.args.train_batch_size,
|
||||
sampler=train_sampler,
|
||||
sampler=train_sampler if not isinstance(self.train_dataset, IterableDataset) else None,
|
||||
collate_fn=self.data_collator.collate_batch,
|
||||
)
|
||||
|
||||
@@ -259,9 +267,10 @@ class Trainer:
|
||||
else:
|
||||
sampler = SequentialSampler(eval_dataset)
|
||||
|
||||
# torch.util.data.Dataset supports samplers, torch.util.data.IterableDataset does not.
|
||||
data_loader = DataLoader(
|
||||
eval_dataset,
|
||||
sampler=sampler,
|
||||
sampler=sampler if not isinstance(eval_dataset, IterableDataset) else None,
|
||||
batch_size=self.args.eval_batch_size,
|
||||
collate_fn=self.data_collator.collate_batch,
|
||||
)
|
||||
@@ -453,13 +462,22 @@ class Trainer:
|
||||
if isinstance(train_dataloader, DataLoader) and isinstance(train_dataloader.sampler, DistributedSampler):
|
||||
train_dataloader.sampler.set_epoch(epoch)
|
||||
|
||||
if t_total - self.global_step < self.num_examples(train_dataloader):
|
||||
total_epoch_steps = min(t_total - self.global_step, len(train_dataloader))
|
||||
else:
|
||||
total_epoch_steps = None
|
||||
|
||||
if is_tpu_available():
|
||||
parallel_loader = pl.ParallelLoader(train_dataloader, [self.args.device]).per_device_loader(
|
||||
self.args.device
|
||||
)
|
||||
epoch_iterator = tqdm(parallel_loader, desc="Iteration", disable=not self.is_local_master())
|
||||
epoch_iterator = tqdm(
|
||||
parallel_loader, desc="Iteration", disable=not self.is_local_master(), total=total_epoch_steps
|
||||
)
|
||||
else:
|
||||
epoch_iterator = tqdm(train_dataloader, desc="Iteration", disable=not self.is_local_master())
|
||||
epoch_iterator = tqdm(
|
||||
train_dataloader, desc="Iteration", disable=not self.is_local_master(), total=total_epoch_steps
|
||||
)
|
||||
|
||||
for step, inputs in enumerate(epoch_iterator):
|
||||
|
||||
@@ -635,7 +653,7 @@ class Trainer:
|
||||
logger.info("Saving model checkpoint to %s", output_dir)
|
||||
# Save a trained model and configuration using `save_pretrained()`.
|
||||
# They can then be reloaded using `from_pretrained()`
|
||||
if not isinstance(self.model, PreTrainedModel):
|
||||
if not isinstance(self.model, PreTrainedModel) and not isinstance(self.model, nn.Module):
|
||||
raise ValueError("Trainer.model appears to not be a PreTrainedModel")
|
||||
self.model.save_pretrained(output_dir)
|
||||
|
||||
@@ -739,14 +757,18 @@ class Trainer:
|
||||
logger.info(" Num examples = %d", self.num_examples(dataloader))
|
||||
logger.info(" Batch size = %d", batch_size)
|
||||
eval_losses: List[float] = []
|
||||
preds: torch.Tensor = None
|
||||
label_ids: torch.Tensor = None
|
||||
evaluation_values: Dict[str, torch.Tensor] = {}
|
||||
model.eval()
|
||||
|
||||
if is_tpu_available():
|
||||
dataloader = pl.ParallelLoader(dataloader, [self.args.device]).per_device_loader(self.args.device)
|
||||
|
||||
for inputs in tqdm(dataloader, desc=description):
|
||||
evaluation_total_steps = 0
|
||||
evaluation_iterator = tqdm(
|
||||
dataloader, desc=description, total=self.args.max_eval_steps if self.args.max_eval_steps > 0 else None
|
||||
)
|
||||
|
||||
for inputs in evaluation_iterator:
|
||||
has_labels = any(inputs.get(k) is not None for k in ["labels", "lm_labels", "masked_lm_labels"])
|
||||
|
||||
for k, v in inputs.items():
|
||||
@@ -754,46 +776,76 @@ class Trainer:
|
||||
|
||||
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 not prediction_loss_only:
|
||||
if preds is None:
|
||||
preds = logits.detach()
|
||||
else:
|
||||
preds = torch.cat((preds, logits.detach()), dim=0)
|
||||
if inputs.get("labels") is not None:
|
||||
if label_ids is None:
|
||||
label_ids = inputs["labels"].detach()
|
||||
if self.manage_evaluation_predictions is None:
|
||||
if has_labels:
|
||||
step_eval_loss, logits = outputs[:2]
|
||||
eval_losses += [step_eval_loss.mean().item()]
|
||||
else:
|
||||
label_ids = torch.cat((label_ids, inputs["labels"].detach()), dim=0)
|
||||
logits = outputs[0]
|
||||
|
||||
if not prediction_loss_only:
|
||||
if "preds" not in evaluation_values:
|
||||
evaluation_values["preds"] = logits.detach()
|
||||
else:
|
||||
evaluation_values["preds"] = torch.cat(
|
||||
(evaluation_values["preds"], logits.detach()), dim=0
|
||||
)
|
||||
if inputs.get("labels") is not None:
|
||||
if "label_ids" not in evaluation_values:
|
||||
evaluation_values["label_ids"] = inputs["labels"].detach()
|
||||
else:
|
||||
evaluation_values["label_ids"] = torch.cat(
|
||||
(evaluation_values["label_ids"], inputs["labels"].detach()), dim=0
|
||||
)
|
||||
else:
|
||||
if not len(evaluation_values):
|
||||
evaluation_values = self.manage_evaluation_predictions(outputs)
|
||||
else:
|
||||
for key, value in self.manage_evaluation_predictions(outputs).items():
|
||||
evaluation_values[key] = torch.cat((evaluation_values[key], value), dim=0)
|
||||
|
||||
evaluation_total_steps += 1
|
||||
|
||||
if evaluation_total_steps > self.args.max_eval_steps > 0:
|
||||
evaluation_iterator.close()
|
||||
break
|
||||
|
||||
if self.args.local_rank != -1:
|
||||
# In distributed mode, concatenate all results from all nodes:
|
||||
if preds is not None:
|
||||
preds = self.distributed_concat(preds, num_total_examples=self.num_examples(dataloader))
|
||||
if label_ids is not None:
|
||||
label_ids = self.distributed_concat(label_ids, num_total_examples=self.num_examples(dataloader))
|
||||
total_examples = (
|
||||
self.args.max_eval_steps if self.args.max_eval_steps > 0 else self.num_examples(dataloader)
|
||||
)
|
||||
for key, value in evaluation_values.items():
|
||||
if value is not None:
|
||||
evaluation_values[key] = self.distributed_concat(value, num_total_examples=total_examples)
|
||||
elif is_tpu_available():
|
||||
# tpu-comment: Get all predictions and labels from all worker shards of eval dataset
|
||||
if preds is not None:
|
||||
preds = xm.mesh_reduce("eval_preds", preds, torch.cat)
|
||||
if label_ids is not None:
|
||||
label_ids = xm.mesh_reduce("eval_label_ids", label_ids, torch.cat)
|
||||
for key, value in evaluation_values.items():
|
||||
if value is not None:
|
||||
evaluation_values[key] = xm.mesh_reduce(key, value, torch.cat)
|
||||
|
||||
# Finally, turn the aggregated tensors into numpy arrays.
|
||||
if preds is not None:
|
||||
preds = preds.cpu().numpy()
|
||||
if label_ids is not None:
|
||||
label_ids = label_ids.cpu().numpy()
|
||||
evaluation_values_numpy: Dict[str, np.ndarray] = {}
|
||||
for key, value in evaluation_values.items():
|
||||
if value is not None:
|
||||
evaluation_values_numpy[key] = value.cpu().numpy()
|
||||
|
||||
if self.compute_metrics is not None and preds is not None and label_ids is not None:
|
||||
metrics = self.compute_metrics(EvalPrediction(predictions=preds, label_ids=label_ids))
|
||||
if self.compute_metrics is not None:
|
||||
if self.manage_evaluation_predictions is not None and self.eval_prediction_mapping is not None:
|
||||
eval_predictions = self.eval_prediction_mapping(evaluation_values_numpy)
|
||||
metrics = self.compute_metrics(eval_predictions)
|
||||
elif "preds" in evaluation_values_numpy and "label_ids" in evaluation_values_numpy:
|
||||
eval_predictions = EvalPrediction(
|
||||
predictions=evaluation_values_numpy["preds"], label_ids=evaluation_values_numpy["label_ids"]
|
||||
)
|
||||
metrics = self.compute_metrics(eval_predictions)
|
||||
else:
|
||||
metrics = {}
|
||||
eval_predictions = EvalPrediction(np.array(0), np.array(0))
|
||||
else:
|
||||
metrics = {}
|
||||
eval_predictions = EvalPrediction(np.array(0), np.array(0))
|
||||
if len(eval_losses) > 0:
|
||||
metrics["eval_loss"] = np.mean(eval_losses)
|
||||
|
||||
@@ -802,7 +854,9 @@ class Trainer:
|
||||
if not key.startswith("eval_"):
|
||||
metrics[f"eval_{key}"] = metrics.pop(key)
|
||||
|
||||
return PredictionOutput(predictions=preds, label_ids=label_ids, metrics=metrics)
|
||||
return PredictionOutput(
|
||||
predictions=eval_predictions.predictions, label_ids=eval_predictions.label_ids, metrics=metrics
|
||||
)
|
||||
|
||||
def distributed_concat(self, tensor: torch.Tensor, num_total_examples: int) -> torch.Tensor:
|
||||
assert self.args.local_rank != -1
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import Dict, NamedTuple, Optional
|
||||
from typing import Dict, NamedTuple, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
@@ -9,8 +9,8 @@ class EvalPrediction(NamedTuple):
|
||||
to compute metrics.
|
||||
"""
|
||||
|
||||
predictions: np.ndarray
|
||||
label_ids: np.ndarray
|
||||
predictions: Union[np.ndarray, Dict[str, np.ndarray]]
|
||||
label_ids: Union[np.ndarray, Dict[str, np.ndarray]]
|
||||
|
||||
|
||||
class PredictionOutput(NamedTuple):
|
||||
|
||||
@@ -95,6 +95,9 @@ class TrainingArguments:
|
||||
default=-1,
|
||||
metadata={"help": "If > 0: set total number of training steps to perform. Override num_train_epochs."},
|
||||
)
|
||||
max_eval_steps: int = field(
|
||||
default=-1, metadata={"help": "If > 0: set total number of eval steps to perform."},
|
||||
)
|
||||
warmup_steps: int = field(default=0, metadata={"help": "Linear warmup over warmup_steps."})
|
||||
|
||||
logging_dir: Optional[str] = field(default=None, metadata={"help": "Tensorboard log dir."})
|
||||
|
||||
Reference in New Issue
Block a user