Compare commits
19
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4114a96831 | ||
|
|
fa6a4414bb | ||
|
|
1b073be34e | ||
|
|
4abf78f2b1 | ||
|
|
bf5819f6f4 | ||
|
|
8fd0275066 | ||
|
|
daf2a4a37e | ||
|
|
8202330d24 | ||
|
|
995c47b1fe | ||
|
|
e36a51ed58 | ||
|
|
d0be398f50 | ||
|
|
896f8aaefc | ||
|
|
520b558a6a | ||
|
|
4ee2f6f4d8 | ||
|
|
0ca151b168 | ||
|
|
6fc3849dce | ||
|
|
ba4c3a9a07 | ||
|
|
342db2aaf6 | ||
|
|
1108ec9bfd |
@@ -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:]
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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))
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user