Compare commits

...
Author SHA1 Message Date
LysandreJik 4114a96831 Test model parallelization 2020-11-10 20:03:35 -05:00
alexorona fa6a4414bb Update modeling_t5.py 2020-11-05 23:51:19 -08:00
alexorona 1b073be34e Update model_parallel_utils.py 2020-11-05 23:40:15 -08:00
alexorona 4abf78f2b1 Update trainer.py 2020-11-05 23:38:55 -08:00
alexorona bf5819f6f4 Update modeling_t5.py 2020-11-05 23:38:19 -08:00
alexorona 8fd0275066 Update modeling_gpt2.py 2020-11-05 23:36:17 -08:00
alexorona daf2a4a37e Update on modeling_t5.py 2020-10-18 10:37:58 -07:00
alexorona 8202330d24 Reformatted modeling_t5 for code quality check
Note: parellelism not yet introduced into t5. Just doing this to get past checks.
2020-10-18 10:27:36 -07:00
alexorona 995c47b1fe Minor changes and reverses t5 commit. 2020-10-18 10:10:33 -07:00
alexorona e36a51ed58 Update modeling_t5.py 2020-10-18 09:58:46 -07:00
alexorona d0be398f50 Update model_parallel_utils.py 2020-10-16 22:05:49 -07:00
alexorona 896f8aaefc Update trainer.py 2020-10-16 22:05:05 -07:00
alexorona 520b558a6a Update training_args.py 2020-10-16 22:03:44 -07:00
alexorona 4ee2f6f4d8 Update modeling_gpt2.py
Fixed a bug when no device_map was provided.
2020-10-16 22:02:22 -07:00
alexorona 0ca151b168 Update training_args.py 2020-10-15 19:31:26 -07:00
alexorona 6fc3849dce Update model_parallel_utils.py 2020-10-15 19:29:29 -07:00
alexorona ba4c3a9a07 Update trainer.py 2020-10-15 19:28:29 -07:00
alexorona 342db2aaf6 Update modeling_gpt2.py 2020-10-15 19:26:55 -07:00
alexorona 1108ec9bfd Added gpt2 model parallelism 2020-10-13 23:02:04 -07:00
8 changed files with 962 additions and 150 deletions
+239 -30
View File
@@ -33,7 +33,11 @@ from .file_utils import (
add_start_docstrings_to_callable,
replace_return_docstrings,
)
from .modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast, SequenceClassifierOutputWithPast
from .modeling_outputs import (
BaseModelOutputWithPast,
CausalLMOutputWithPast,
SequenceClassifierOutputWithPast,
)
from .modeling_utils import (
Conv1D,
PreTrainedModel,
@@ -42,6 +46,7 @@ from .modeling_utils import (
prune_conv1d_layer,
)
from .utils import logging
from .utils.model_parallel_utils import assert_device_map, get_device_map
logger = logging.get_logger(__name__)
@@ -124,7 +129,10 @@ class Attention(nn.Module):
# [switch nx => n_state from Block to Attention to keep identical to TF implem]
assert n_state % config.n_head == 0
self.register_buffer(
"bias", torch.tril(torch.ones((n_ctx, n_ctx), dtype=torch.uint8)).view(1, 1, n_ctx, n_ctx)
"bias",
torch.tril(torch.ones((n_ctx, n_ctx), dtype=torch.uint8)).view(
1, 1, n_ctx, n_ctx
),
)
self.register_buffer("masked_bias", torch.tensor(-1e4))
self.n_head = config.n_head
@@ -147,7 +155,9 @@ class Attention(nn.Module):
heads, index = find_pruneable_heads_and_indices(
heads, self.n_head, self.split_size // self.n_head, self.pruned_heads
)
index_attn = torch.cat([index, index + self.split_size, index + (2 * self.split_size)])
index_attn = torch.cat(
[index, index + self.split_size, index + (2 * self.split_size)]
)
# Prune conv1d layers
self.c_attn = prune_conv1d_layer(self.c_attn, index_attn, dim=1)
@@ -158,7 +168,9 @@ class Attention(nn.Module):
self.n_head = self.n_head - len(heads)
self.pruned_heads = self.pruned_heads.union(heads)
def _attn(self, q, k, v, attention_mask=None, head_mask=None, output_attentions=False):
def _attn(
self, q, k, v, attention_mask=None, head_mask=None, output_attentions=False
):
w = torch.matmul(q, k)
if self.scale:
w = w / (float(v.size(-1)) ** 0.5)
@@ -214,7 +226,9 @@ class Attention(nn.Module):
self, "q_attn"
), "If class is used as cross attention, the weights `q_attn` have to be defined. Please make sure to instantiate class with `Attention(..., is_cross_attention=True)`."
query = self.q_attn(hidden_states)
key, value = self.c_attn(encoder_hidden_states).split(self.split_size, dim=2)
key, value = self.c_attn(encoder_hidden_states).split(
self.split_size, dim=2
)
attention_mask = encoder_attention_mask
else:
query, key, value = self.c_attn(hidden_states).split(self.split_size, dim=2)
@@ -223,16 +237,23 @@ class Attention(nn.Module):
key = self.split_heads(key, k=True)
value = self.split_heads(value)
if layer_past is not None:
past_key, past_value = layer_past[0].transpose(-2, -1), layer_past[1] # transpose back cf below
past_key, past_value = (
layer_past[0].transpose(-2, -1),
layer_past[1],
) # transpose back cf below
key = torch.cat((past_key, key), dim=-1)
value = torch.cat((past_value, value), dim=-2)
if use_cache is True:
present = torch.stack((key.transpose(-2, -1), value)) # transpose to have same shapes for stacking
present = torch.stack(
(key.transpose(-2, -1), value)
) # transpose to have same shapes for stacking
else:
present = (None,)
attn_outputs = self._attn(query, key, value, attention_mask, head_mask, output_attentions)
attn_outputs = self._attn(
query, key, value, attention_mask, head_mask, output_attentions
)
a = attn_outputs[0]
a = self.merge_heads(a)
@@ -267,8 +288,12 @@ class Block(nn.Module):
self.attn = Attention(hidden_size, n_ctx, config, scale)
self.ln_2 = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)
if config.add_cross_attention:
self.crossattention = Attention(hidden_size, n_ctx, config, scale, is_cross_attention=True)
self.ln_cross_attn = nn.LayerNorm(hidden_size, eps=config.layer_norm_epsilon)
self.crossattention = Attention(
hidden_size, n_ctx, config, scale, is_cross_attention=True
)
self.ln_cross_attn = nn.LayerNorm(
hidden_size, eps=config.layer_norm_epsilon
)
self.mlp = MLP(inner_dim, config)
def forward(
@@ -311,7 +336,9 @@ class Block(nn.Module):
attn_output = cross_attn_outputs[0]
# residual connection
hidden_states = hidden_states + attn_output
outputs = outputs + cross_attn_outputs[1:] # add cross attentions if we output attention weights
outputs = (
outputs + cross_attn_outputs[1:]
) # add cross attentions if we output attention weights
feed_forward_hidden_states = self.mlp(self.ln_2(hidden_states))
# residual connection
@@ -472,6 +499,48 @@ GPT2_INPUTS_DOCSTRING = r"""
Whether or not to return a :class:`~transformers.file_utils.ModelOutput` instead of a plain tuple.
"""
PARALLELIZE_DOCSTRING = r"""
Uses a device map to distribute attention modules of the model across several devices. If no device map is given, it
will evenly distribute blocks across all devices.
Args:
device_map (:obj:`Dict[int, list]`, optional, defaults to None):
A dictionary that maps attention modules to devices. Note that the embedding module and LMHead are
always automatically mapped to the first device (for esoteric reasons). That means that the first
device should have fewer attention modules mapped to it than other devices.
For reference, the gpt2 models have the following number of attention modules:
- gpt2: 12
- gpt2-medium: 24
- gpt2-large: 36
- gpt2-xl: 48
Example::
Here is an example of a device map on a machine with 4 GPUs using gpt2-xl, which has a total of 48 attention modules:
model = GPT2LMHeadModel.from_pretrained('gpt2-xl')
device_map = {0: [0, 1, 2, 3, 4, 5, 6, 7, 8],
1: [9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21],
2: [22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34],
3: [35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47]}
model.parallelize(device_map)
"""
DEPARALLELIZE_DOCSTRING = r"""
Moves the model to cpu from a model parallel state.
Example::
On a 4 GPU machine with gpt2-large:
model = GPT2LMHeadModel.from_pretrained('gpt2-large')
device_map = {0: [0, 1, 2, 3, 4, 5, 6, 7],
1: [8, 9, 10, 11, 12, 13, 14, 15],
2: [16, 17, 18, 19, 20, 21, 22, 23],
3: [24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35]}
model.parallelize(device_map) # Splits the model across several devices
model.deparallelize() # Put the model back on cpu and cleans memory by calling torch.cuda.empty_cache()
"""
@add_start_docstrings(
"The bare GPT2 Model transformer outputting raw hidden-states without any specific head on top.",
@@ -484,11 +553,58 @@ class GPT2Model(GPT2PreTrainedModel):
self.wte = nn.Embedding(config.vocab_size, config.n_embd)
self.wpe = nn.Embedding(config.n_positions, config.n_embd)
self.drop = nn.Dropout(config.embd_pdrop)
self.h = nn.ModuleList([Block(config.n_ctx, config, scale=True) for _ in range(config.n_layer)])
self.h = nn.ModuleList(
[Block(config.n_ctx, config, scale=True) for _ in range(config.n_layer)]
)
self.ln_f = nn.LayerNorm(config.n_embd, eps=config.layer_norm_epsilon)
self.init_weights()
# Model parallel
self.model_parallel = False
self.device_map = None
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def parallelize(self, device_map=None):
# Check validity of device_map
self.device_map = (
get_device_map(len(self.h), range(torch.cuda.device_count()))
if device_map is None
else device_map
)
assert_device_map(self.device_map, len(self.h))
self.model_parallel = True
self.first_device = (
"cpu"
if "cpu" in self.device_map.keys()
else "cuda:" + str(min(self.device_map.keys()))
)
self.last_device = "cuda:" + str(max(self.device_map.keys()))
self.wte = self.wte.to(self.first_device)
self.wpe = self.wpe.to(self.first_device)
# Load onto devices
for k, v in self.device_map.items():
for block in v:
cuda_device = "cuda:" + str(k)
self.h[block] = self.h[block].to(cuda_device)
# ln_f to last
self.ln_f = self.ln_f.to(self.last_device)
@add_start_docstrings(DEPARALLELIZE_DOCSTRING)
def deparallelize(self):
self.model_parallel = False
self.device_map = None
self.first_device = "cpu"
self.last_device = "cpu"
self.wte = self.wte.to("cpu")
self.wpe = self.wpe.to("cpu")
for index in range(len(self.h)):
self.h[index] = self.h[index].to("cpu")
self.ln_f = self.ln_f.to("cpu")
torch.cuda.empty_cache()
def get_input_embeddings(self):
return self.wte
@@ -534,15 +650,25 @@ class GPT2Model(GPT2PreTrainedModel):
past_key_values = kwargs.pop("past")
assert kwargs == {}, f"Unexpected keyword arguments: {list(kwargs.keys())}."
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_attentions = (
output_attentions
if output_attentions is not None
else self.config.output_attentions
)
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
output_hidden_states
if output_hidden_states is not None
else self.config.output_hidden_states
)
use_cache = use_cache if use_cache is not None else self.config.use_cache
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return_dict = (
return_dict if return_dict is not None else self.config.use_return_dict
)
if input_ids is not None and inputs_embeds is not None:
raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
raise ValueError(
"You cannot specify both input_ids and inputs_embeds at the same time"
)
elif input_ids is not None:
input_shape = input_ids.size()
input_ids = input_ids.view(-1, input_shape[-1])
@@ -565,7 +691,12 @@ class GPT2Model(GPT2PreTrainedModel):
past_length = past_key_values[0][0].size(-2)
if position_ids is None:
device = input_ids.device if input_ids is not None else inputs_embeds.device
position_ids = torch.arange(past_length, input_shape[-1] + past_length, dtype=torch.long, device=device)
position_ids = torch.arange(
past_length,
input_shape[-1] + past_length,
dtype=torch.long,
device=device,
)
position_ids = position_ids.unsqueeze(0).view(-1, input_shape[-1])
# Attention mask.
@@ -590,7 +721,11 @@ class GPT2Model(GPT2PreTrainedModel):
# If a 2D ou 3D attention mask is provided for the cross-attention
# we need to make broadcastabe to [batch_size, num_heads, seq_length, seq_length]
if self.config.add_cross_attention and encoder_hidden_states is not None:
encoder_batch_size, encoder_sequence_length, _ = encoder_hidden_states.size()
(
encoder_batch_size,
encoder_sequence_length,
_,
) = encoder_hidden_states.size()
encoder_hidden_shape = (encoder_batch_size, encoder_sequence_length)
if encoder_attention_mask is None:
encoder_attention_mask = torch.ones(encoder_hidden_shape, device=device)
@@ -620,15 +755,33 @@ class GPT2Model(GPT2PreTrainedModel):
all_attentions = () if output_attentions else None
all_hidden_states = () if output_hidden_states else None
for i, (block, layer_past) in enumerate(zip(self.h, past_key_values)):
# Model parallel
if self.model_parallel:
torch.cuda.set_device(hidden_states.device)
# Ensure layer_past is on same device as hidden_states (might not be correct)
if layer_past is not None:
layer_past = layer_past.to(hidden_states.device)
# Ensure that attention_mask is always on the same device as hidden_states
if attention_mask is not None:
attention_mask = attention_mask.to(hidden_states.device)
if isinstance(head_mask, torch.Tensor):
head_mask = head_mask.to(hidden_states.device)
if output_hidden_states:
all_hidden_states = all_hidden_states + (hidden_states.view(*output_shape),)
all_hidden_states = all_hidden_states + (
hidden_states.view(*output_shape),
)
if getattr(self.config, "gradient_checkpointing", False):
def create_custom_forward(module):
def custom_forward(*inputs):
# checkpointing only works with tuple returns, not with lists
return tuple(output for output in module(*inputs, use_cache, output_attentions))
return tuple(
output
for output in module(*inputs, use_cache, output_attentions)
)
return custom_forward
@@ -660,6 +813,12 @@ class GPT2Model(GPT2PreTrainedModel):
if output_attentions:
all_attentions = all_attentions + (outputs[2],)
# Model Parallel: If it's the last layer for that device, put things on the next device
if self.model_parallel:
for k, v in self.device_map.items():
if i == v[-1] and "cuda:" + str(k) != self.last_device:
hidden_states = hidden_states.to("cuda:" + str(k + 1))
hidden_states = self.ln_f(hidden_states)
hidden_states = hidden_states.view(*output_shape)
@@ -668,7 +827,11 @@ class GPT2Model(GPT2PreTrainedModel):
all_hidden_states = all_hidden_states + (hidden_states,)
if not return_dict:
return tuple(v for v in [hidden_states, presents, all_hidden_states, all_attentions] if v is not None)
return tuple(
v
for v in [hidden_states, presents, all_hidden_states, all_attentions]
if v is not None
)
return BaseModelOutputWithPast(
last_hidden_state=hidden_states,
@@ -693,6 +856,29 @@ class GPT2LMHeadModel(GPT2PreTrainedModel):
self.init_weights()
self.model_parallel = False
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def parallelize(self, device_map=None):
self.device_map = (
get_device_map(len(self.transformer.h), range(torch.cuda.device_count()))
if device_map is None
else device_map
)
assert_device_map(self.device_map, len(self.transformer.h))
self.transformer.parallelize(self.device_map)
self.lm_head = self.lm_head.to(self.transformer.first_device)
self.model_parallel = True
@add_start_docstrings(DEPARALLELIZE_DOCSTRING)
def deparallelize(self):
self.transformer.deparallelize()
self.transformer = self.transformer.to("cpu")
self.lm_head = self.lm_head.to("cpu")
self.model_parallel = False
torch.cuda.empty_cache()
def get_output_embeddings(self):
return self.lm_head
@@ -747,7 +933,9 @@ class GPT2LMHeadModel(GPT2PreTrainedModel):
)
past_key_values = kwargs.pop("past")
assert kwargs == {}, f"Unexpected keyword arguments: {list(kwargs.keys())}."
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return_dict = (
return_dict if return_dict is not None else self.config.use_return_dict
)
transformer_outputs = self.transformer(
input_ids,
@@ -766,6 +954,11 @@ class GPT2LMHeadModel(GPT2PreTrainedModel):
)
hidden_states = transformer_outputs[0]
# Set device for model parallelism
if self.model_parallel:
torch.cuda.set_device(self.transformer.first_device)
hidden_states = hidden_states.to(self.lm_head.weight.device)
lm_logits = self.lm_head(hidden_states)
loss = None
@@ -775,7 +968,9 @@ class GPT2LMHeadModel(GPT2PreTrainedModel):
shift_labels = labels[..., 1:].contiguous()
# Flatten the tokens
loss_fct = CrossEntropyLoss()
loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
loss = loss_fct(
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)
)
if not return_dict:
output = (lm_logits,) + transformer_outputs[1:]
@@ -823,7 +1018,9 @@ class GPT2DoubleHeadsModel(GPT2PreTrainedModel):
}
@add_start_docstrings_to_callable(GPT2_INPUTS_DOCSTRING)
@replace_return_docstrings(output_type=GPT2DoubleHeadsModelOutput, config_class=_CONFIG_FOR_DOC)
@replace_return_docstrings(
output_type=GPT2DoubleHeadsModelOutput, config_class=_CONFIG_FOR_DOC
)
def forward(
self,
input_ids=None,
@@ -899,7 +1096,9 @@ class GPT2DoubleHeadsModel(GPT2PreTrainedModel):
)
past_key_values = kwargs.pop("past")
assert kwargs == {}, f"Unexpected keyword arguments: {list(kwargs.keys())}."
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return_dict = (
return_dict if return_dict is not None else self.config.use_return_dict
)
transformer_outputs = self.transformer(
input_ids,
@@ -923,13 +1122,17 @@ class GPT2DoubleHeadsModel(GPT2PreTrainedModel):
mc_loss = None
if mc_labels is not None:
loss_fct = CrossEntropyLoss()
mc_loss = loss_fct(mc_logits.view(-1, mc_logits.size(-1)), mc_labels.view(-1))
mc_loss = loss_fct(
mc_logits.view(-1, mc_logits.size(-1)), mc_labels.view(-1)
)
lm_loss = None
if labels is not None:
shift_logits = lm_logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fct = CrossEntropyLoss()
lm_loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
lm_loss = loss_fct(
shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)
)
if not return_dict:
output = (lm_logits, mc_logits) + transformer_outputs[1:]
@@ -1003,7 +1206,9 @@ class GPT2ForSequenceClassification(GPT2PreTrainedModel):
If :obj:`config.num_labels == 1` a regression loss is computed (Mean-Square loss),
If :obj:`config.num_labels > 1` a classification loss is computed (Cross-Entropy).
"""
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return_dict = (
return_dict if return_dict is not None else self.config.use_return_dict
)
transformer_outputs = self.transformer(
input_ids,
@@ -1033,7 +1238,9 @@ class GPT2ForSequenceClassification(GPT2PreTrainedModel):
sequence_lengths = -1
else:
if input_ids is not None:
sequence_lengths = torch.ne(input_ids, self.config.pad_token_id).sum(-1) - 1
sequence_lengths = (
torch.ne(input_ids, self.config.pad_token_id).sum(-1) - 1
)
else:
sequence_lengths = -1
logger.warning(
@@ -1051,7 +1258,9 @@ class GPT2ForSequenceClassification(GPT2PreTrainedModel):
loss = loss_fct(pooled_logits.view(-1), labels.view(-1))
else:
loss_fct = CrossEntropyLoss()
loss = loss_fct(pooled_logits.view(-1, self.num_labels), labels.view(-1))
loss = loss_fct(
pooled_logits.view(-1, self.num_labels), labels.view(-1)
)
if not return_dict:
output = (pooled_logits,) + transformer_outputs[1:]
+192 -7
View File
@@ -36,6 +36,7 @@ from .file_utils import (
from .modeling_outputs import BaseModelOutput, BaseModelOutputWithPast, Seq2SeqLMOutput, Seq2SeqModelOutput
from .modeling_utils import PreTrainedModel, find_pruneable_heads_and_indices, prune_linear_layer
from .utils import logging
from .utils.model_parallel_utils import assert_device_map, get_device_map
logger = logging.get_logger(__name__)
@@ -151,7 +152,48 @@ def load_tf_weights_in_t5(model, config, tf_checkpoint_path):
# - torch.nn.Module for the layers and
# - PreTrainedModel for the models (it-self a sub-class of torch.nn.Module)
####################################################
PARALLELIZE_DOCSTRING = r"""
Uses a device map to distribute attention modules of the model across several devices. If no device map is given, it
will evenly distribute blocks across all devices.
Args:
device_map (:obj:`Dict[int, list]`, optional, defaults to None):
A dictionary that maps attention modules to devices. Note that the embedding module and LMHead are
always automatically mapped to the first device (for esoteric reasons). That means that the first
device should have fewer attention modules mapped to it than other devices.
For reference, the t5 models have the following number of attention modules:
- t5-small: 6
- t5-base: 12
- t5-large: 24
- t5-3b: 24
- t5-11b: 24
Example::
Here is an example of a device map on a machine with 4 GPUs using t5-3b, which has a total of 24 attention modules:
model = T5ForConditionalGeneration.from_pretrained('t5-3b')
device_map = {0: [0, 1, 2],
1: [3, 4, 5, 6, 7, 8, 9],
2: [10, 11, 12, 13, 14, 15, 16],
3: [17, 18, 19, 20, 21, 22, 23]}
model.parallelize(device_map)
"""
DEPARALLELIZE_DOCSTRING = r"""
Moves the model to cpu from a model parallel state.
Example::
On a 4 GPU machine with t5-3b:
model = T5ForConditionalGeneration.from_pretrained('t5-3b')
device_map = {0: [0, 1, 2],
1: [3, 4, 5, 6, 7, 8, 9],
2: [10, 11, 12, 13, 14, 15, 16],
3: [17, 18, 19, 20, 21, 22, 23]}
model.parallelize(device_map) # Splits the model across several devices
model.deparallelize() # Put the model back on cpu and cleans memory by calling torch.cuda.empty_cache()
"""
class T5LayerNorm(nn.Module):
def __init__(self, hidden_size, eps=1e-6):
@@ -661,6 +703,43 @@ class T5Stack(T5PreTrainedModel):
self.init_weights()
# Model parallel
self.model_parallel = False
self.device_map = None
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def parallelize(self, device_map=None):
# Check validity of device_map
self.device_map = get_device_map(len(self.block), torch.cuda.device_count()) if device_map is None else device_map
assert_device_map(self.device_map, len(self.block))
self.model_parallel = True
self.first_device = "cpu" if "cpu" in self.device_map.keys() else "cuda:" + str(min(self.device_map.keys()))
self.last_device = "cuda:" + str(max(self.device_map.keys()))
# Load onto devices
for k, v in self.device_map.items():
for layer in v:
cuda_device = "cuda:" + str(k)
self.block[layer] = self.block[layer].to(cuda_device)
# Set embed_tokens to first layer
self.embed_tokens = self.embed_tokens.to(self.first_device)
# Set final layer norm to last device
self.final_layer_norm = self.final_layer_norm.to(self.last_device)
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def deparallelize(self):
self.model_parallel = False
self.device_map = None
self.first_device = "cpu"
self.last_device = "cpu"
for i in range(len(self.block)):
self.block[i] = self.block[i].to("cpu")
self.embed_tokens = self.embed_tokens.to("cpu")
self.final_layer_norm = self.final_layer_norm.to("cpu")
torch.cuda.empty_cache()
def get_input_embeddings(self):
return self.embed_tokens
@@ -684,15 +763,19 @@ class T5Stack(T5PreTrainedModel):
output_hidden_states=None,
return_dict=None,
):
# # Model parallel
if self.model_parallel:
torch.cuda.set_device(self.first_device)
self.embed_tokens = self.embed_tokens.to(self.first_device)
use_cache = use_cache if use_cache is not None else self.config.use_cache
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if input_ids is not None and inputs_embeds is not None:
err_msg_prefix = "decoder_" if self.is_decoder else ""
raise ValueError(
f"You cannot specify both {err_msg_prefix}inputs and {err_msg_prefix}inputs_embeds at the same time"
@@ -705,7 +788,6 @@ class T5Stack(T5PreTrainedModel):
else:
err_msg_prefix = "decoder_" if self.is_decoder else ""
raise ValueError(f"You have to specify either {err_msg_prefix}inputs or {err_msg_prefix}inputs_embeds")
if inputs_embeds is None:
assert self.embed_tokens is not None, "You have to intialize the model with valid token embeddings"
inputs_embeds = self.embed_tokens(input_ids)
@@ -719,7 +801,6 @@ class T5Stack(T5PreTrainedModel):
assert self.is_decoder, ":obj:`use_cache` can only be set to `True` if {} is used as a decoder".format(
self
)
if attention_mask is None:
attention_mask = torch.ones(batch_size, mask_seq_length).to(inputs_embeds.device)
if self.is_decoder and encoder_attention_mask is None and encoder_hidden_states is not None:
@@ -727,7 +808,6 @@ class T5Stack(T5PreTrainedModel):
encoder_attention_mask = torch.ones(
batch_size, encoder_seq_length, device=inputs_embeds.device, dtype=torch.long
)
# initialize past_key_values with `None` if past does not exist
if past_key_values is None:
past_key_values = [None] * len(self.block)
@@ -751,6 +831,21 @@ class T5Stack(T5PreTrainedModel):
hidden_states = self.dropout(inputs_embeds)
for i, (layer_module, past_key_value) in enumerate(zip(self.block, past_key_values)):
# Model parallel
if self.model_parallel:
torch.cuda.set_device(hidden_states.device)
# Ensure that attention_mask is always on the same device as hidden_states
if attention_mask is not None:
attention_mask = attention_mask.to(hidden_states.device)
if position_bias is not None:
position_bias = position_bias.to(hidden_states.device)
if encoder_hidden_states is not None:
encoder_hidden_states = encoder_hidden_states.to(hidden_states.device)
if encoder_extended_attention_mask is not None:
encoder_extended_attention_mask = encoder_extended_attention_mask.to(hidden_states.device)
if encoder_decoder_position_bias is not None:
encoder_decoder_position_bias = encoder_decoder_position_bias.to(hidden_states.device)
if output_hidden_states:
all_hidden_states = all_hidden_states + (hidden_states,)
@@ -782,6 +877,11 @@ class T5Stack(T5PreTrainedModel):
if output_attentions:
all_attentions = all_attentions + (layer_outputs[2],) # We keep only self-attention weights for now
# Model Parallel: If it's the last layer for that device, put things on the next device
if self.model_parallel:
for k, v in self.device_map.items():
if i == v[-1] and "cuda:" + str(k) != self.last_device:
hidden_states = hidden_states.to("cuda:" + str(k + 1))
hidden_states = self.final_layer_norm(hidden_states)
hidden_states = self.dropout(hidden_states)
@@ -904,7 +1004,6 @@ T5_INPUTS_DOCSTRING = r"""
Whether or not to return a :class:`~transformers.file_utils.ModelOutput` instead of a plain tuple.
"""
@add_start_docstrings(
"The bare T5 Model transformer outputting raw hidden-states" "without any specific head on top.",
T5_START_DOCSTRING,
@@ -927,6 +1026,32 @@ class T5Model(T5PreTrainedModel):
self.init_weights()
# Model parallel
self.model_parallel = False
self.device_map = None
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def parallelize(self, device_map=None):
self.device_map = (
get_device_map(len(self.encoder.block), range(torch.cuda.device_count())) if device_map is None else device_map
)
assert_device_map(self.device_map, len(self.encoder.block))
self.encoder.parallelize(self.device_map)
self.decoder.parallelize(self.device_map)
self.model_parallel = True
@add_start_docstrings(DEPARALLELIZE_DOCSTRING)
def deparallelize(self):
self.encoder.deparallelize()
self.decoder.deparallelize()
self.encoder = self.encoder.to("cpu")
self.decoder = self.decoder.to("cpu")
self.model_parallel = False
self.device_map = None
torch.cuda.empty_cache()
def get_input_embeddings(self):
return self.shared
@@ -1020,6 +1145,19 @@ class T5Model(T5PreTrainedModel):
)
hidden_states = encoder_outputs[0]
if self.model_parallel:
torch.cuda.set_device(self.decoder.first_device)
# Set device for model parallelism
if self.model_parallel:
torch.cuda.set_device(self.decoder.first_device)
hidden_states = hidden_states.to(self.decoder.first_device)
if decoder_input_ids is not None:
decoder_input_ids = decoder_input_ids.to(self.decoder.first_device)
if attention_mask is not None:
attention_mask = attention_mask.to(self.decoder.first_device)
if decoder_attention_mask is not None:
decoder_attention_mask = decoder_attention_mask.to(self.decoder.first_device)
# Decode
decoder_outputs = self.decoder(
@@ -1075,6 +1213,34 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
self.init_weights()
# Model parallel
self.model_parallel = False
self.device_map = None
@add_start_docstrings(PARALLELIZE_DOCSTRING)
def parallelize(self, device_map=None):
self.device_map = (
get_device_map(len(self.encoder.block), range(torch.cuda.device_count())) if device_map is None else device_map
)
assert_device_map(self.device_map, len(self.encoder.block))
self.encoder.parallelize(self.device_map)
self.decoder.parallelize(self.device_map)
self.lm_head = self.lm_head.to(self.decoder.first_device)
self.model_parallel = True
@add_start_docstrings(DEPARALLELIZE_DOCSTRING)
def deparallelize(self):
self.encoder.deparallelize()
self.decoder.deparallelize()
self.encoder = self.encoder.to("cpu")
self.decoder = self.decoder.to("cpu")
self.lm_head = self.lm_head.to("cpu")
self.model_parallel = False
self.device_map = None
torch.cuda.empty_cache()
def get_input_embeddings(self):
return self.shared
@@ -1139,7 +1305,6 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
>>> input_ids = tokenizer("summarize: studies have shown that owning a dog is good for you ", return_tensors="pt").input_ids # Batch size 1
>>> outputs = model.generate(input_ids)
"""
if "lm_labels" in kwargs:
warnings.warn(
"The `lm_labels` argument is deprecated and will be removed in a future version, use `labels` instead.",
@@ -1175,6 +1340,7 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
elif return_dict and not isinstance(encoder_outputs, BaseModelOutput):
encoder_outputs = BaseModelOutput(
last_hidden_state=encoder_outputs[0],
@@ -1184,6 +1350,9 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
hidden_states = encoder_outputs[0]
if self.model_parallel:
torch.cuda.set_device(self.decoder.first_device)
if labels is not None and decoder_input_ids is None and decoder_inputs_embeds is None:
# get decoder inputs from shifting lm labels to the right
decoder_input_ids = self._shift_right(labels)
@@ -1197,6 +1366,17 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
if decoder_inputs_embeds is not None:
decoder_inputs_embeds = decoder_inputs_embeds[:, -1:]
# Set device for model parallelism
if self.model_parallel:
torch.cuda.set_device(self.decoder.first_device)
hidden_states = hidden_states.to(self.decoder.first_device)
if decoder_input_ids is not None:
decoder_input_ids = decoder_input_ids.to(self.decoder.first_device)
if attention_mask is not None:
attention_mask = attention_mask.to(self.decoder.first_device)
if decoder_attention_mask is not None:
decoder_attention_mask = decoder_attention_mask.to(self.decoder.first_device)
# Decode
decoder_outputs = self.decoder(
input_ids=decoder_input_ids,
@@ -1213,6 +1393,11 @@ class T5ForConditionalGeneration(T5PreTrainedModel):
)
sequence_output = decoder_outputs[0]
# Set device for model parallelism
if self.model_parallel:
torch.cuda.set_device(self.encoder.first_device)
self.lm_head = self.lm_head.to(self.encoder.first_device)
sequence_output = sequence_output.to(self.lm_head.weight.device)
# Rescale output before projecting on vocab
# See https://github.com/tensorflow/mesh/blob/fa19d69eafc9a482aff0b59ddd96b025c0cb207d/mesh_tensorflow/transformer/transformer.py#L586
sequence_output = sequence_output * (self.model_dim ** -0.5)
+370 -111
View File
@@ -33,7 +33,11 @@ from torch.utils.data.dataset import Dataset
from torch.utils.data.distributed import DistributedSampler
from torch.utils.data.sampler import RandomSampler, SequentialSampler
from .data.data_collator import DataCollator, DataCollatorWithPadding, default_data_collator
from .data.data_collator import (
DataCollator,
DataCollatorWithPadding,
default_data_collator,
)
from .file_utils import WEIGHTS_NAME, is_datasets_available, is_torch_tpu_available
from .integrations import (
default_hp_search_backend,
@@ -203,11 +207,16 @@ class Trainer:
model_init: Callable[[], PreTrainedModel] = None,
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None,
callbacks: Optional[List[TrainerCallback]] = None,
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (
None,
None,
),
**kwargs,
):
if args is None:
logger.info("No `TrainingArguments` passed, using the current path as `output_dir`.")
logger.info(
"No `TrainingArguments` passed, using the current path as `output_dir`."
)
args = TrainingArguments("tmp_trainer")
self.args = args
# Seed must be set before instantiating the model when using model
@@ -216,25 +225,45 @@ class Trainer:
model is not None or model_init is not None
), "You must provide a model to use `Trainer`, either by using the `model` argument or the `model_init` argument."
self.model_init = model_init
if model is None and model_init is not None:
model = self.call_model_init()
self.model = model.to(args.device) if model is not None else None
default_collator = default_data_collator if tokenizer is None else DataCollatorWithPadding(tokenizer)
self.data_collator = data_collator if data_collator is not None else default_collator
# Model parallel
self.model = model if model else None
if not self.args.model_parallel and self.model is not None:
self.model = self.model.to(args.device)
default_collator = (
default_data_collator
if tokenizer is None
else DataCollatorWithPadding(tokenizer)
)
self.data_collator = (
data_collator if data_collator is not None else default_collator
)
self.train_dataset = train_dataset
self.eval_dataset = eval_dataset
self.tokenizer = tokenizer
self.compute_metrics = compute_metrics
self.optimizer, self.lr_scheduler = optimizers
if model_init is not None and (self.optimizer is not None or self.lr_scheduler is not None):
if model_init is not None and (
self.optimizer is not None or self.lr_scheduler is not None
):
raise RuntimeError(
"Passing a `model_init` is incompatible with providing the `optimizers` argument."
"You should subclass `Trainer` and override the `create_optimizer_and_scheduler` method."
)
callbacks = DEFAULT_CALLBACKS if callbacks is None else DEFAULT_CALLBACKS + callbacks
self.callback_handler = CallbackHandler(callbacks, self.model, self.optimizer, self.lr_scheduler)
self.add_callback(PrinterCallback if self.args.disable_tqdm else ProgressCallback)
callbacks = (
DEFAULT_CALLBACKS if callbacks is None else DEFAULT_CALLBACKS + callbacks
)
self.callback_handler = CallbackHandler(
callbacks, self.model, self.optimizer, self.lr_scheduler
)
self.add_callback(
PrinterCallback if self.args.disable_tqdm else ProgressCallback
)
# Deprecated arguments
if "tb_writer" in kwargs:
@@ -267,7 +296,9 @@ class Trainer:
# Set an xla_device flag on the model's config.
# We'll find a more elegant and not need to do this in the future.
self.model.config.xla_device = True
if not callable(self.data_collator) and callable(getattr(self.data_collator, "collate_batch", None)):
if not callable(self.data_collator) and callable(
getattr(self.data_collator, "collate_batch", None)
):
self.data_collator = self.data_collator.collate_batch
warnings.warn(
(
@@ -297,8 +328,14 @@ class Trainer:
if type(self.model) in MODEL_FOR_QUESTION_ANSWERING_MAPPING.values()
else ["labels"]
)
self.label_names = default_label_names if self.args.label_names is None else self.args.label_names
self.control = self.callback_handler.on_init_end(self.args, self.state, self.control)
self.label_names = (
default_label_names
if self.args.label_names is None
else self.args.label_names
)
self.control = self.callback_handler.on_init_end(
self.args, self.state, self.control
)
def add_callback(self, callback):
"""
@@ -338,7 +375,9 @@ class Trainer:
"""
self.callback_handler.remove_callback(callback)
def _remove_unused_columns(self, dataset: "datasets.Dataset", description: Optional[str] = None):
def _remove_unused_columns(
self, dataset: "datasets.Dataset", description: Optional[str] = None
):
if not self.args.remove_unused_columns:
return
# Inspect model forward signature to keep only the arguments it accepts.
@@ -388,11 +427,15 @@ class Trainer:
num_workers=self.args.dataloader_num_workers,
)
def _get_eval_sampler(self, eval_dataset: Dataset) -> Optional[torch.utils.data.sampler.Sampler]:
def _get_eval_sampler(
self, eval_dataset: Dataset
) -> Optional[torch.utils.data.sampler.Sampler]:
if isinstance(eval_dataset, torch.utils.data.IterableDataset):
return None
elif is_torch_tpu_available():
return SequentialDistributedSampler(eval_dataset, num_replicas=xm.xrt_world_size(), rank=xm.get_ordinal())
return SequentialDistributedSampler(
eval_dataset, num_replicas=xm.xrt_world_size(), rank=xm.get_ordinal()
)
elif self.args.local_rank != -1:
return SequentialDistributedSampler(eval_dataset)
else:
@@ -414,7 +457,11 @@ class Trainer:
"""
if eval_dataset is None and self.eval_dataset is None:
raise ValueError("Trainer: evaluation requires an eval_dataset.")
elif eval_dataset is not None and is_datasets_available() and isinstance(eval_dataset, datasets.Dataset):
elif (
eval_dataset is not None
and is_datasets_available()
and isinstance(eval_dataset, datasets.Dataset)
):
self._remove_unused_columns(eval_dataset, description="evaluation")
eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset
eval_sampler = self._get_eval_sampler(eval_dataset)
@@ -466,11 +513,19 @@ class Trainer:
no_decay = ["bias", "LayerNorm.weight"]
optimizer_grouped_parameters = [
{
"params": [p for n, p in self.model.named_parameters() if not any(nd in n for nd in no_decay)],
"params": [
p
for n, p in self.model.named_parameters()
if not any(nd in n for nd in no_decay)
],
"weight_decay": self.args.weight_decay,
},
{
"params": [p for n, p in self.model.named_parameters() if any(nd in n for nd in no_decay)],
"params": [
p
for n, p in self.model.named_parameters()
if any(nd in n for nd in no_decay)
],
"weight_decay": 0.0,
},
]
@@ -482,7 +537,9 @@ class Trainer:
)
if self.lr_scheduler is None:
self.lr_scheduler = get_linear_schedule_with_warmup(
self.optimizer, num_warmup_steps=self.args.warmup_steps, num_training_steps=num_training_steps
self.optimizer,
num_warmup_steps=self.args.warmup_steps,
num_training_steps=num_training_steps,
)
def num_examples(self, dataloader: DataLoader) -> int:
@@ -495,7 +552,11 @@ class Trainer:
""" HP search setup code """
if self.hp_search_backend is None or trial is None:
return
params = self.hp_space(trial) if self.hp_search_backend == HPSearchBackend.OPTUNA else trial
params = (
self.hp_space(trial)
if self.hp_search_backend == HPSearchBackend.OPTUNA
else trial
)
for key, value in params.items():
if not hasattr(self.args, key):
raise AttributeError(
@@ -510,7 +571,10 @@ class Trainer:
logger.info("Trial:", trial.params)
def _report_to_hp_search(
self, trial: Union["optuna.Trial", Dict[str, Any]], epoch: int, metrics: Dict[str, float]
self,
trial: Union["optuna.Trial", Dict[str, Any]],
epoch: int,
metrics: Dict[str, float],
):
if self.hp_search_backend is None or trial is None:
return
@@ -529,12 +593,21 @@ class Trainer:
return
with tune.checkpoint_dir(step=self.state.global_step) as checkpoint_dir:
self.args.output_dir = checkpoint_dir
output_dir = os.path.join(self.args.output_dir, f"{PREFIX_CHECKPOINT_DIR}-{self.state.global_step}")
output_dir = os.path.join(
self.args.output_dir,
f"{PREFIX_CHECKPOINT_DIR}-{self.state.global_step}",
)
self.save_model(output_dir)
if self.is_world_master():
self.state.save_to_json(os.path.join(output_dir, "trainer_state.json"))
torch.save(self.optimizer.state_dict(), os.path.join(output_dir, "optimizer.pt"))
torch.save(self.lr_scheduler.state_dict(), os.path.join(output_dir, "scheduler.pt"))
torch.save(
self.optimizer.state_dict(),
os.path.join(output_dir, "optimizer.pt"),
)
torch.save(
self.lr_scheduler.state_dict(),
os.path.join(output_dir, "scheduler.pt"),
)
def call_model_init(self, trial=None):
model_init_argcount = len(inspect.signature(self.model_init).parameters)
@@ -547,7 +620,11 @@ class Trainer:
return model
def train(self, model_path: Optional[str] = None, trial: Union["optuna.Trial", Dict[str, Any]] = None):
def train(
self,
model_path: Optional[str] = None,
trial: Union["optuna.Trial", Dict[str, Any]] = None,
):
"""
Main training entry point.
@@ -568,14 +645,18 @@ class Trainer:
model = self.call_model_init(trial)
self.model = model.to(self.args.device)
# Model parallel
if not self.args.model_parallel:
self.model = model.to(self.args.device)
# Reinitializes optimizer and scheduler
self.optimizer, self.lr_scheduler = None, None
# Data loader and number of training steps
train_dataloader = self.get_train_dataloader()
num_update_steps_per_epoch = len(train_dataloader) // self.args.gradient_accumulation_steps
num_update_steps_per_epoch = (
len(train_dataloader) // self.args.gradient_accumulation_steps
)
num_update_steps_per_epoch = max(num_update_steps_per_epoch, 1)
if self.args.max_steps > 0:
max_steps = self.args.max_steps
@@ -598,21 +679,30 @@ class Trainer:
):
# Load in optimizer and scheduler states
self.optimizer.load_state_dict(
torch.load(os.path.join(model_path, "optimizer.pt"), map_location=self.args.device)
torch.load(
os.path.join(model_path, "optimizer.pt"),
map_location=self.args.device,
)
)
with warnings.catch_warnings(record=True) as caught_warnings:
self.lr_scheduler.load_state_dict(torch.load(os.path.join(model_path, "scheduler.pt")))
self.lr_scheduler.load_state_dict(
torch.load(os.path.join(model_path, "scheduler.pt"))
)
reissue_pt_warnings(caught_warnings)
# Mixed precision training with apex (torch < 1.6)
model = self.model
if self.args.fp16 and _use_apex:
if not is_apex_available():
raise ImportError("Please install apex from https://www.github.com/nvidia/apex to use fp16 training.")
model, self.optimizer = amp.initialize(model, self.optimizer, opt_level=self.args.fp16_opt_level)
raise ImportError(
"Please install apex from https://www.github.com/nvidia/apex to use fp16 training."
)
model, self.optimizer = amp.initialize(
model, self.optimizer, opt_level=self.args.fp16_opt_level
)
# Multi-gpu training (should be after apex fp16 initialization)
if self.args.n_gpu > 1:
if self.args.n_gpu > 1 and not self.args.model_parallel:
model = torch.nn.DataParallel(model)
# Distributed training (should be after apex fp16 initialization)
@@ -637,14 +727,26 @@ class Trainer:
total_train_batch_size = (
self.args.train_batch_size
* self.args.gradient_accumulation_steps
* (torch.distributed.get_world_size() if self.args.local_rank != -1 else 1)
* (
torch.distributed.get_world_size()
if self.args.local_rank != -1
else 1
)
)
logger.info("***** Running training *****")
logger.info(" Num examples = %d", self.num_examples(train_dataloader))
logger.info(" Num Epochs = %d", num_train_epochs)
logger.info(" Instantaneous batch size per device = %d", self.args.per_device_train_batch_size)
logger.info(" Total train batch size (w. parallel, distributed & accumulation) = %d", total_train_batch_size)
logger.info(" Gradient Accumulation steps = %d", self.args.gradient_accumulation_steps)
logger.info(
" Instantaneous batch size per device = %d",
self.args.per_device_train_batch_size,
)
logger.info(
" Total train batch size (w. parallel, distributed & accumulation) = %d",
total_train_batch_size,
)
logger.info(
" Gradient Accumulation steps = %d", self.args.gradient_accumulation_steps
)
logger.info(" Total optimization steps = %d", max_steps)
self.state.epoch = 0
@@ -652,15 +754,28 @@ class Trainer:
steps_trained_in_current_epoch = 0
# Check if continuing training from a checkpoint
if model_path and os.path.isfile(os.path.join(model_path, "trainer_state.json")):
self.state = TrainerState.load_from_json(os.path.join(model_path, "trainer_state.json"))
if model_path and os.path.isfile(
os.path.join(model_path, "trainer_state.json")
):
self.state = TrainerState.load_from_json(
os.path.join(model_path, "trainer_state.json")
)
epochs_trained = self.state.global_step // num_update_steps_per_epoch
steps_trained_in_current_epoch = self.state.global_step % (num_update_steps_per_epoch)
steps_trained_in_current_epoch = self.state.global_step % (
num_update_steps_per_epoch
)
logger.info(" Continuing training from checkpoint, will skip to saved global_step")
logger.info(
" Continuing training from checkpoint, will skip to saved global_step"
)
logger.info(" Continuing training from epoch %d", epochs_trained)
logger.info(" Continuing training from global step %d", self.state.global_step)
logger.info(" Will skip the first %d steps in the first epoch", steps_trained_in_current_epoch)
logger.info(
" Continuing training from global step %d", self.state.global_step
)
logger.info(
" Will skip the first %d steps in the first epoch",
steps_trained_in_current_epoch,
)
# Update the references
self.callback_handler.model = self.model
@@ -679,16 +794,20 @@ class Trainer:
self._total_flos = self.state.total_flos
model.zero_grad()
self.control = self.callback_handler.on_train_begin(self.args, self.state, self.control)
self.control = self.callback_handler.on_train_begin(
self.args, self.state, self.control
)
for epoch in range(epochs_trained, num_train_epochs):
if isinstance(train_dataloader, DataLoader) and isinstance(train_dataloader.sampler, DistributedSampler):
if isinstance(train_dataloader, DataLoader) and isinstance(
train_dataloader.sampler, DistributedSampler
):
train_dataloader.sampler.set_epoch(epoch)
if is_torch_tpu_available():
parallel_loader = pl.ParallelLoader(train_dataloader, [self.args.device]).per_device_loader(
self.args.device
)
parallel_loader = pl.ParallelLoader(
train_dataloader, [self.args.device]
).per_device_loader(self.args.device)
epoch_iterator = parallel_loader
else:
epoch_iterator = train_dataloader
@@ -697,7 +816,9 @@ class Trainer:
if self.args.past_index >= 0:
self._past = None
self.control = self.callback_handler.on_epoch_begin(self.args, self.state, self.control)
self.control = self.callback_handler.on_epoch_begin(
self.args, self.state, self.control
)
for step, inputs in enumerate(epoch_iterator):
@@ -707,7 +828,9 @@ class Trainer:
continue
if (step + 1) % self.args.gradient_accumulation_steps == 0:
self.control = self.callback_handler.on_step_begin(self.args, self.state, self.control)
self.control = self.callback_handler.on_step_begin(
self.args, self.state, self.control
)
if (
((step + 1) % self.args.gradient_accumulation_steps != 0)
@@ -727,11 +850,17 @@ class Trainer:
):
if self.args.fp16 and _use_native_amp:
self.scaler.unscale_(self.optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), self.args.max_grad_norm)
torch.nn.utils.clip_grad_norm_(
model.parameters(), self.args.max_grad_norm
)
elif self.args.fp16 and _use_apex:
torch.nn.utils.clip_grad_norm_(amp.master_params(self.optimizer), self.args.max_grad_norm)
torch.nn.utils.clip_grad_norm_(
amp.master_params(self.optimizer), self.args.max_grad_norm
)
else:
torch.nn.utils.clip_grad_norm_(model.parameters(), self.args.max_grad_norm)
torch.nn.utils.clip_grad_norm_(
model.parameters(), self.args.max_grad_norm
)
if is_torch_tpu_available():
xm.optimizer_step(self.optimizer)
@@ -745,14 +874,18 @@ class Trainer:
model.zero_grad()
self.state.global_step += 1
self.state.epoch = epoch + (step + 1) / len(epoch_iterator)
self.control = self.callback_handler.on_step_end(self.args, self.state, self.control)
self.control = self.callback_handler.on_step_end(
self.args, self.state, self.control
)
self._maybe_log_save_evalute(tr_loss, model, trial, epoch)
if self.control.should_epoch_stop or self.control.should_training_stop:
break
self.control = self.callback_handler.on_epoch_end(self.args, self.state, self.control)
self.control = self.callback_handler.on_epoch_end(
self.args, self.state, self.control
)
self._maybe_log_save_evalute(tr_loss, model, trial, epoch)
if self.args.tpu_metrics_debug or self.args.debug:
@@ -771,27 +904,42 @@ class Trainer:
# Clean the state at the end of training
delattr(self, "_past")
logger.info("\n\nTraining completed. Do not forget to share your model on huggingface.co/models =)\n\n")
if self.args.load_best_model_at_end and self.state.best_model_checkpoint is not None:
logger.info(
"\n\nTraining completed. Do not forget to share your model on huggingface.co/models =)\n\n"
)
if (
self.args.load_best_model_at_end
and self.state.best_model_checkpoint is not None
):
logger.info(
f"Loading best model from {self.state.best_model_checkpoint} (score: {self.state.best_metric})."
)
if isinstance(model, PreTrainedModel):
self.model = model.from_pretrained(self.state.best_model_checkpoint)
self.model = self.model.to(self.args.device)
if not self.args.model_parallel:
self.model = model.to(self.args.device)
else:
state_dict = torch.load(os.path.join(self.state.best_model_checkpoint, WEIGHTS_NAME))
state_dict = torch.load(
os.path.join(self.state.best_model_checkpoint, WEIGHTS_NAME)
)
self.model.load_state_dict(state_dict)
self.control = self.callback_handler.on_train_end(self.args, self.state, self.control)
self.control = self.callback_handler.on_train_end(
self.args, self.state, self.control
)
return TrainOutput(self.state.global_step, tr_loss.item() / self.state.global_step)
return TrainOutput(
self.state.global_step, tr_loss.item() / self.state.global_step
)
def _maybe_log_save_evalute(self, tr_loss, model, trial, epoch):
if self.control.should_log:
logs: Dict[str, float] = {}
tr_loss_scalar = tr_loss.item()
logs["loss"] = (tr_loss_scalar - self._logging_loss_scalar) / self.args.logging_steps
logs["loss"] = (
tr_loss_scalar - self._logging_loss_scalar
) / self.args.logging_steps
# backward compatibility for pytorch schedulers
logs["learning_rate"] = (
self.lr_scheduler.get_last_lr()[0]
@@ -806,23 +954,35 @@ class Trainer:
if self.control.should_evaluate:
metrics = self.evaluate()
self._report_to_hp_search(trial, epoch, metrics)
self.control = self.callback_handler.on_evaluate(self.args, self.state, self.control, metrics)
self.control = self.callback_handler.on_evaluate(
self.args, self.state, self.control, metrics
)
if self.control.should_save:
self._save_checkpoint(model, trial, metrics=metrics)
self.control = self.callback_handler.on_save(self.args, self.state, self.control)
self.control = self.callback_handler.on_save(
self.args, self.state, self.control
)
def _save_checkpoint(self, model, trial, metrics=None):
# In all cases (even distributed/parallel), self.model is always a reference
# to the model we want to save.
if hasattr(model, "module"):
assert model.module is self.model, f"Module {model.module} should be a reference to self.model"
assert (
model.module is self.model
), f"Module {model.module} should be a reference to self.model"
else:
assert model is self.model, f"Model {model} should be a reference to self.model"
assert (
model is self.model
), f"Model {model} should be a reference to self.model"
# Save model checkpoint
checkpoint_folder = f"{PREFIX_CHECKPOINT_DIR}-{self.state.global_step}"
if self.hp_search_backend is not None and trial is not None:
run_id = trial.number if self.hp_search_backend == HPSearchBackend.OPTUNA else tune.get_trial_id()
run_id = (
trial.number
if self.hp_search_backend == HPSearchBackend.OPTUNA
else tune.get_trial_id()
)
checkpoint_folder += f"-run-{run_id}"
output_dir = os.path.join(self.args.output_dir, checkpoint_folder)
@@ -832,14 +992,24 @@ class Trainer:
# Save optimizer and scheduler
if is_torch_tpu_available():
xm.rendezvous("saving_optimizer_states")
xm.save(self.optimizer.state_dict(), os.path.join(output_dir, "optimizer.pt"))
xm.save(
self.optimizer.state_dict(), os.path.join(output_dir, "optimizer.pt")
)
with warnings.catch_warnings(record=True) as caught_warnings:
xm.save(self.lr_scheduler.state_dict(), os.path.join(output_dir, "scheduler.pt"))
xm.save(
self.lr_scheduler.state_dict(),
os.path.join(output_dir, "scheduler.pt"),
)
reissue_pt_warnings(caught_warnings)
elif self.is_world_process_zero():
torch.save(self.optimizer.state_dict(), os.path.join(output_dir, "optimizer.pt"))
torch.save(
self.optimizer.state_dict(), os.path.join(output_dir, "optimizer.pt")
)
with warnings.catch_warnings(record=True) as caught_warnings:
torch.save(self.lr_scheduler.state_dict(), os.path.join(output_dir, "scheduler.pt"))
torch.save(
self.lr_scheduler.state_dict(),
os.path.join(output_dir, "scheduler.pt"),
)
reissue_pt_warnings(caught_warnings)
# Determine the new best metric / best model checkpoint
@@ -873,7 +1043,7 @@ class Trainer:
n_trials: int = 20,
direction: str = "minimize",
backend: Optional[Union["str", HPSearchBackend]] = None,
**kwargs
**kwargs,
) -> BestRun:
"""
Launch an hyperparameter search using ``optuna`` or ``Ray Tune``. The optimized quantity is determined by
@@ -924,7 +1094,9 @@ class Trainer:
)
backend = HPSearchBackend(backend)
if backend == HPSearchBackend.OPTUNA and not is_optuna_available():
raise RuntimeError("You picked the optuna backend, but it is not installed. Use `pip install optuna`.")
raise RuntimeError(
"You picked the optuna backend, but it is not installed. Use `pip install optuna`."
)
if backend == HPSearchBackend.RAY and not is_ray_available():
raise RuntimeError(
"You picked the Ray Tune backend, but it is not installed. Use `pip install 'ray[tune]'`."
@@ -937,9 +1109,17 @@ class Trainer:
)
self.hp_space = default_hp_space[backend] if hp_space is None else hp_space
self.compute_objective = default_compute_objective if compute_objective is None else compute_objective
self.compute_objective = (
default_compute_objective
if compute_objective is None
else compute_objective
)
run_hp_search = run_hp_search_optuna if backend == HPSearchBackend.OPTUNA else run_hp_search_ray
run_hp_search = (
run_hp_search_optuna
if backend == HPSearchBackend.OPTUNA
else run_hp_search_ray
)
best_run = run_hp_search(self, n_trials, direction, **kwargs)
self.hp_search_backend = None
@@ -967,11 +1147,15 @@ class Trainer:
if self._total_flos is not None:
self.store_flos()
logs["total_flos"] = self.state.total_flos
self.control = self.callback_handler.on_log(self.args, self.state, self.control, logs)
self.control = self.callback_handler.on_log(
self.args, self.state, self.control, logs
)
output = {**logs, **{"step": self.state.global_step}}
self.state.log_history.append(output)
def _prepare_inputs(self, inputs: Dict[str, Union[torch.Tensor, Any]]) -> Dict[str, Union[torch.Tensor, Any]]:
def _prepare_inputs(
self, inputs: Dict[str, Union[torch.Tensor, Any]]
) -> Dict[str, Union[torch.Tensor, Any]]:
"""
Prepare :obj:`inputs` before feeding them to the model, converting them to tensors if they are not already and
handling potential state.
@@ -985,7 +1169,9 @@ class Trainer:
return inputs
def training_step(self, model: nn.Module, inputs: Dict[str, Union[torch.Tensor, Any]]) -> torch.Tensor:
def training_step(
self, model: nn.Module, inputs: Dict[str, Union[torch.Tensor, Any]]
) -> torch.Tensor:
"""
Perform a training step on a batch of inputs.
@@ -1019,7 +1205,7 @@ class Trainer:
else:
loss = self.compute_loss(model, inputs)
if self.args.n_gpu > 1:
if self.args.n_gpu > 1 and not self.args.model_parallel:
loss = loss.mean() # mean() to average on multi-gpu parallel training
if self.args.gradient_accumulation_steps > 1:
@@ -1057,7 +1243,10 @@ class Trainer:
This method is deprecated, use :meth:`~transformers.Trainer.is_local_process_zero` instead.
"""
warnings.warn("This method is deprecated, use `Trainer.is_local_process_zero()` instead.", FutureWarning)
warnings.warn(
"This method is deprecated, use `Trainer.is_local_process_zero()` instead.",
FutureWarning,
)
return self.is_local_process_zero()
def is_local_process_zero(self) -> bool:
@@ -1079,7 +1268,10 @@ class Trainer:
This method is deprecated, use :meth:`~transformers.Trainer.is_world_process_zero` instead.
"""
warnings.warn("This method is deprecated, use `Trainer.is_world_process_zero()` instead.", FutureWarning)
warnings.warn(
"This method is deprecated, use `Trainer.is_world_process_zero()` instead.",
FutureWarning,
)
return self.is_world_process_zero()
def is_world_process_zero(self) -> bool:
@@ -1116,7 +1308,9 @@ class Trainer:
# They can then be reloaded using `from_pretrained()`
xm.rendezvous("saving_checkpoint")
if not isinstance(self.model, PreTrainedModel):
logger.info("Trainer.model is not a `PreTrainedModel`, only saving its state dict.")
logger.info(
"Trainer.model is not a `PreTrainedModel`, only saving its state dict."
)
state_dict = self.model.state_dict()
xm.save(state_dict, os.path.join(output_dir, WEIGHTS_NAME))
else:
@@ -1131,7 +1325,9 @@ class Trainer:
# Save a trained model and configuration using `save_pretrained()`.
# They can then be reloaded using `from_pretrained()`
if not isinstance(self.model, PreTrainedModel):
logger.info("Trainer.model is not a `PreTrainedModel`, only saving its state dict.")
logger.info(
"Trainer.model is not a `PreTrainedModel`, only saving its state dict."
)
state_dict = self.model.state_dict()
torch.save(state_dict, os.path.join(output_dir, WEIGHTS_NAME))
else:
@@ -1146,14 +1342,20 @@ class Trainer:
# Storing the number of floating-point operations that went into the model
if self._total_flos is not None:
if self.args.local_rank != -1:
self.state.total_flos = distributed_broadcast_scalars([self._total_flos]).sum().item()
self.state.total_flos = (
distributed_broadcast_scalars([self._total_flos]).sum().item()
)
else:
self.state.total_flos = self._total_flos
def _sorted_checkpoints(self, checkpoint_prefix=PREFIX_CHECKPOINT_DIR, use_mtime=False) -> List[str]:
def _sorted_checkpoints(
self, checkpoint_prefix=PREFIX_CHECKPOINT_DIR, use_mtime=False
) -> List[str]:
ordering_and_checkpoint_path = []
glob_checkpoints = [str(x) for x in Path(self.args.output_dir).glob(f"{checkpoint_prefix}-*")]
glob_checkpoints = [
str(x) for x in Path(self.args.output_dir).glob(f"{checkpoint_prefix}-*")
]
for path in glob_checkpoints:
if use_mtime:
@@ -1161,14 +1363,21 @@ class Trainer:
else:
regex_match = re.match(f".*{checkpoint_prefix}-([0-9]+)", path)
if regex_match and regex_match.groups():
ordering_and_checkpoint_path.append((int(regex_match.groups()[0]), path))
ordering_and_checkpoint_path.append(
(int(regex_match.groups()[0]), path)
)
checkpoints_sorted = sorted(ordering_and_checkpoint_path)
checkpoints_sorted = [checkpoint[1] for checkpoint in checkpoints_sorted]
# Make sure we don't delete the best model.
if self.state.best_model_checkpoint is not None:
best_model_index = checkpoints_sorted.index(self.state.best_model_checkpoint)
checkpoints_sorted[best_model_index], checkpoints_sorted[best_model_index][-1] = (
best_model_index = checkpoints_sorted.index(
self.state.best_model_checkpoint
)
(
checkpoints_sorted[best_model_index],
checkpoints_sorted[best_model_index][-1],
) = (
checkpoints_sorted[-1],
checkpoints_sorted[best_model_index],
)
@@ -1183,10 +1392,16 @@ class Trainer:
if len(checkpoints_sorted) <= self.args.save_total_limit:
return
number_of_checkpoints_to_delete = max(0, len(checkpoints_sorted) - self.args.save_total_limit)
number_of_checkpoints_to_delete = max(
0, len(checkpoints_sorted) - self.args.save_total_limit
)
checkpoints_to_be_deleted = checkpoints_sorted[:number_of_checkpoints_to_delete]
for checkpoint in checkpoints_to_be_deleted:
logger.info("Deleting older checkpoint [{}] due to args.save_total_limit".format(checkpoint))
logger.info(
"Deleting older checkpoint [{}] due to args.save_total_limit".format(
checkpoint
)
)
shutil.rmtree(checkpoint)
def evaluate(self, eval_dataset: Optional[Dataset] = None) -> Dict[str, float]:
@@ -1244,7 +1459,10 @@ class Trainer:
return self.prediction_loop(test_dataloader, description="Prediction")
def prediction_loop(
self, dataloader: DataLoader, description: str, prediction_loss_only: Optional[bool] = None
self,
dataloader: DataLoader,
description: str,
prediction_loss_only: Optional[bool] = None,
) -> PredictionOutput:
"""
Prediction/evaluation loop, shared by :obj:`Trainer.evaluate()` and :obj:`Trainer.predict()`.
@@ -1256,15 +1474,19 @@ class Trainer:
"The `_prediction_loop` method is deprecated and won't be called in a future version, define `prediction_loop` in your subclass.",
FutureWarning,
)
return self._prediction_loop(dataloader, description, prediction_loss_only=prediction_loss_only)
return self._prediction_loop(
dataloader, description, prediction_loss_only=prediction_loss_only
)
prediction_loss_only = (
prediction_loss_only if prediction_loss_only is not None else self.args.prediction_loss_only
prediction_loss_only
if prediction_loss_only is not None
else self.args.prediction_loss_only
)
model = self.model
# multi-gpu eval
if self.args.n_gpu > 1:
# multi-gpu eval without model parallel
if self.args.n_gpu > 1 and not self.args.model_parallel:
model = torch.nn.DataParallel(model)
else:
model = self.model
@@ -1281,7 +1503,9 @@ class Trainer:
model.eval()
if is_torch_tpu_available():
dataloader = pl.ParallelLoader(dataloader, [self.args.device]).per_device_loader(self.args.device)
dataloader = pl.ParallelLoader(
dataloader, [self.args.device]
).per_device_loader(self.args.device)
if self.args.past_index >= 0:
self._past = None
@@ -1289,15 +1513,23 @@ class Trainer:
self.callback_handler.eval_dataloader = dataloader
for inputs in dataloader:
loss, logits, labels = self.prediction_step(model, inputs, prediction_loss_only)
loss, logits, labels = self.prediction_step(
model, inputs, prediction_loss_only
)
batch_size = inputs[list(inputs.keys())[0]].shape[0]
if loss is not None:
eval_losses.extend([loss] * batch_size)
if logits is not None:
preds = logits if preds is None else nested_concat(preds, logits, dim=0)
if labels is not None:
label_ids = labels if label_ids is None else nested_concat(label_ids, labels, dim=0)
self.control = self.callback_handler.on_prediction_step(self.args, self.state, self.control)
label_ids = (
labels
if label_ids is None
else nested_concat(label_ids, labels, dim=0)
)
self.control = self.callback_handler.on_prediction_step(
self.args, self.state, self.control
)
if self.args.past_index and hasattr(self, "_past"):
# Clean the state at the end of the evaluation loop
@@ -1306,9 +1538,13 @@ class Trainer:
if self.args.local_rank != -1:
# In distributed mode, concatenate all results from all nodes:
if preds is not None:
preds = distributed_concat(preds, num_total_examples=self.num_examples(dataloader))
preds = distributed_concat(
preds, num_total_examples=self.num_examples(dataloader)
)
if label_ids is not None:
label_ids = distributed_concat(label_ids, num_total_examples=self.num_examples(dataloader))
label_ids = distributed_concat(
label_ids, num_total_examples=self.num_examples(dataloader)
)
elif is_torch_tpu_available():
# tpu-comment: Get all predictions and labels from all worker shards of eval dataset
if preds is not None:
@@ -1316,7 +1552,9 @@ class Trainer:
if label_ids is not None:
label_ids = nested_xla_mesh_reduce(label_ids, "eval_label_ids")
if eval_losses is not None:
eval_losses = xm.mesh_reduce("eval_losses", torch.tensor(eval_losses), torch.cat).tolist()
eval_losses = xm.mesh_reduce(
"eval_losses", torch.tensor(eval_losses), torch.cat
).tolist()
# Finally, turn the aggregated tensors into numpy arrays.
if preds is not None:
@@ -1324,14 +1562,22 @@ class Trainer:
if label_ids is not None:
label_ids = nested_numpify(label_ids)
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
and preds is not None
and label_ids is not None
):
metrics = self.compute_metrics(
EvalPrediction(predictions=preds, label_ids=label_ids)
)
else:
metrics = {}
if len(eval_losses) > 0:
if self.args.local_rank != -1:
metrics["eval_loss"] = (
distributed_broadcast_scalars(eval_losses, num_total_examples=self.num_examples(dataloader))
distributed_broadcast_scalars(
eval_losses, num_total_examples=self.num_examples(dataloader)
)
.mean()
.item()
)
@@ -1346,7 +1592,10 @@ class Trainer:
return PredictionOutput(predictions=preds, label_ids=label_ids, metrics=metrics)
def prediction_step(
self, model: nn.Module, inputs: Dict[str, Union[torch.Tensor, Any]], prediction_loss_only: bool
self,
model: nn.Module,
inputs: Dict[str, Union[torch.Tensor, Any]],
prediction_loss_only: bool,
) -> Tuple[Optional[float], Optional[torch.Tensor], Optional[torch.Tensor]]:
"""
Perform an evaluation step on :obj:`model` using obj:`inputs`.
@@ -1382,9 +1631,13 @@ class Trainer:
# Slicing so we get a tuple even if `outputs` is a `ModelOutput`.
logits = outputs[:]
if self.args.past_index >= 0:
self._past = outputs[self.args.past_index if has_labels else self.args.past_index - 1]
self._past = outputs[
self.args.past_index if has_labels else self.args.past_index - 1
]
# Remove the past from the logits.
logits = logits[: self.args.past_index - 1] + logits[self.args.past_index :]
logits = (
logits[: self.args.past_index - 1] + logits[self.args.past_index :]
)
if prediction_loss_only:
return (loss, None, None)
@@ -1428,7 +1681,11 @@ class Trainer:
@staticmethod
def _actual_model(
model: Union[torch.nn.DataParallel, torch.nn.parallel.DistributedDataParallel, torch.nn.modules.Module]
model: Union[
torch.nn.DataParallel,
torch.nn.parallel.DistributedDataParallel,
torch.nn.modules.Module,
]
) -> torch.nn.modules.Module:
"""
@@ -1439,7 +1696,9 @@ class Trainer:
Returns:
:obj:`torch.nn.modules.Module`: unwrapped module
"""
if isinstance(model, torch.nn.DataParallel) or isinstance(model, torch.nn.parallel.DistributedDataParallel):
if isinstance(model, torch.nn.DataParallel) or isinstance(
model, torch.nn.parallel.DistributedDataParallel
):
model = model.module
else:
model = model
+18 -2
View File
@@ -54,6 +54,8 @@ class TrainingArguments:
:obj:`"no"`.
do_predict (:obj:`bool`, `optional`, defaults to :obj:`False`):
Whether to run predictions on the test set or not.
model_parallel (:obj:`bool`, `optional`, defaults to :obj:`False`):
If there is more than one device, whether to distribute the model's modules across devices.
evaluation_strategy (:obj:`str` or :class:`~transformers.trainer_utils.EvaluationStrategy`, `optional`, defaults to :obj:`"no"`):
The evaluation strategy to adopt during training. Possible values are:
@@ -186,6 +188,12 @@ class TrainingArguments:
do_train: bool = field(default=False, metadata={"help": "Whether to run training."})
do_eval: bool = field(default=None, metadata={"help": "Whether to run eval on the dev set."})
do_predict: bool = field(default=False, metadata={"help": "Whether to run predictions on the test set."})
model_parallel: bool = field(
default=False,
metadata={
"help": "If there are more than one devices, whether to use model parallelism to distribute the model's modules across devices."
},
)
evaluate_during_training: bool = field(
default=None,
metadata={"help": "Run evaluation during training at each logging step."},
@@ -354,7 +362,11 @@ class TrainingArguments:
"version. Using `--per_device_train_batch_size` is preferred."
)
per_device_batch_size = self.per_gpu_train_batch_size or self.per_device_train_batch_size
return per_device_batch_size * max(1, self.n_gpu)
if not self.model_parallel:
train_batch_size = per_device_batch_size * max(1, self.n_gpu)
else:
train_batch_size = per_device_batch_size
return train_batch_size
@property
def eval_batch_size(self) -> int:
@@ -367,7 +379,11 @@ class TrainingArguments:
"version. Using `--per_device_eval_batch_size` is preferred."
)
per_device_batch_size = self.per_gpu_eval_batch_size or self.per_device_eval_batch_size
return per_device_batch_size * max(1, self.n_gpu)
if not self.model_parallel:
eval_batch_size = per_device_batch_size * max(1, self.n_gpu)
else:
eval_batch_size = per_device_batch_size
return eval_batch_size
@cached_property
@torch_required
@@ -0,0 +1,40 @@
# coding=utf-8
from math import ceil
def assert_device_map(device_map, num_blocks):
blocks = list(range(0, num_blocks))
device_map_blocks = [
item for sublist in list(device_map.values()) for item in sublist
]
# Duplicate check
duplicate_blocks = []
for i in device_map_blocks:
if device_map_blocks.count(i) > 1 and i not in duplicate_blocks:
duplicate_blocks.append(i)
# Missing blocks
missing_blocks = [i for i in blocks if i not in device_map_blocks]
extra_blocks = [i for i in device_map_blocks if i not in blocks]
assert len(duplicate_blocks) == 0, (
"Duplicate attention blocks specified in device_map. Attention blocks must be specified to one device. These attention blocks were specified more than once: "
+ str(duplicate_blocks)
)
assert len(missing_blocks) == 0, (
"There are attention blocks for this model that are not specified in the device_map. Add these attention_blocks to a device on the device_map:"
+ str(missing_blocks)
)
assert len(extra_blocks) == 0, (
"The device_map contains more attention blocks than this model has. Remove these from the device_map:"
+ str(extra_blocks)
)
def get_device_map(n_layers: int, devices: list):
"""Returns a dictionary of layers distributed evenly across all devices."""
layers = list(range(n_layers))
n_blocks = int(ceil(n_layers / len(devices)))
layers_list = list(layers[i : i + n_blocks] for i in range(0, n_layers, n_blocks))
return dict(zip(devices, layers_list))
+93
View File
@@ -66,6 +66,7 @@ class ModelTesterMixin:
test_resize_embeddings = True
test_head_masking = True
test_missing_keys = True
test_model_parallel = False
is_encoder_decoder = False
def _prepare_for_class(self, inputs_dict, model_class, return_labels=False):
@@ -986,6 +987,98 @@ class ModelTesterMixin:
with torch.no_grad():
_ = model(**self._prepare_for_class(inputs_dict, model_class))
@require_multigpu
def test_model_parallelization(self):
if not self.test_model_parallel:
pass
import subprocess
def get_current_gpu_memory_use():
run_process = subprocess.Popen('nvidia-smi --query-gpu=memory.used --format=csv,nounits,noheader', shell=True, stdout=subprocess.PIPE)
memory_usage = run_process.stdout.read().decode('utf-8').strip()
per_device_memory = [int(memory) for memory in memory_usage.split('\n')]
return per_device_memory
# Needs a large model to see the difference.
config = self.model_tester.get_large_model_config()
for model_class in self.all_parallelizable_model_classes:
torch.cuda.empty_cache()
# Retrieve initial memory usage (should be close to 0)
initial_memory = get_current_gpu_memory_use()
# Put model on device
model = model_class(config.from_pretrained("gpt2"))
model.to("cuda:0")
# Retrieve the memory after the model is put on the device
memory_after_model_load = get_current_gpu_memory_use()
del model
torch.cuda.empty_cache()
# Retrieve memory after emptied cache (should be close to 0)
empty_cache = get_current_gpu_memory_use()
# The memory use on that device should be higher than it was initially.
self.assertGreater(memory_after_model_load[0], initial_memory[0])
# Spread model layers over multiple devices
model = model_class(config.from_pretrained("gpt2"))
model.parallelize()
memory_after_parallelization = get_current_gpu_memory_use()
# Assert that the memory use on all devices is higher than it was when loaded only on CPU
for n in range(torch.cuda.device_count()):
self.assertGreater(memory_after_parallelization[n], initial_memory[n])
# Assert that the memory use of the first device is lower than it was when the entire model was loaded on it
self.assertLess(memory_after_parallelization[0], memory_after_model_load[0])
# Assert that the memory use of the second device is higher than it was when the entire model was loaded
# on the other device.
self.assertGreater(memory_after_parallelization[1], memory_after_model_load[1])
del model
torch.cuda.empty_cache()
@require_multigpu
def test_model_parallel_equal_results(self):
if not self.test_model_parallel:
pass
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
for model_class in self.all_parallelizable_model_classes:
inputs_dict = self._prepare_for_class(inputs_dict, model_class)
model = model_class(config)
output = model(**inputs_dict)
model.parallelize()
def cast_to_gpu(dictionary):
output = {}
for k, v in dictionary.items():
if isinstance(v, torch.Tensor):
output[k] = v.to('cuda:0')
else:
output[k] = v
return output
parallel_output = model(**cast_to_gpu(inputs_dict))
for value, parallel_value in zip(output, parallel_output):
if isinstance(value, torch.Tensor):
self.assertTrue(torch.allclose(value, parallel_value.to('cpu'), atol=1e-7))
elif isinstance(value, (Tuple, List)):
for value_, parallel_value_ in zip(value, parallel_value):
self.assertTrue(torch.allclose(value_, parallel_value_.to('cpu'), atol=1e-7))
global_rng = random.Random()
+5
View File
@@ -90,6 +90,9 @@ class GPT2ModelTester:
self.eos_token_id = vocab_size - 1
self.pad_token_id = vocab_size - 1
def get_large_model_config(self):
return GPT2Config.from_pretrained("gpt2")
def prepare_config_and_inputs(self, gradient_checkpointing=False):
input_ids = ids_tensor([self.batch_size, self.seq_length], self.vocab_size)
@@ -384,7 +387,9 @@ class GPT2ModelTest(ModelTesterMixin, unittest.TestCase):
else ()
)
all_generative_model_classes = (GPT2LMHeadModel, GPT2DoubleHeadsModel) if is_torch_available() else ()
all_parallelizable_model_classes = (GPT2LMHeadModel,) if is_torch_available() else ()
test_missing_keys = False
test_model_parallel = True
def setUp(self):
self.model_tester = GPT2ModelTester(self)
+5
View File
@@ -86,6 +86,9 @@ class T5ModelTester:
self.scope = None
self.decoder_layers = decoder_layers
def get_large_model_config(self):
return T5Config.from_pretrained("t5-base")
def prepare_config_and_inputs(self):
input_ids = ids_tensor([self.batch_size, self.encoder_seq_length], self.vocab_size)
decoder_input_ids = ids_tensor([self.batch_size, self.decoder_seq_length], self.vocab_size)
@@ -472,9 +475,11 @@ class T5ModelTest(ModelTesterMixin, unittest.TestCase):
all_model_classes = (T5Model, T5ForConditionalGeneration) if is_torch_available() else ()
all_generative_model_classes = (T5ForConditionalGeneration,) if is_torch_available() else ()
all_parallelizable_model_classes = (T5Model, T5ForConditionalGeneration,) if is_torch_available() else ()
test_pruning = False
test_torchscript = True
test_resize_embeddings = False
test_model_parallel = True
is_encoder_decoder = True
def setUp(self):