Compare commits

...
Author SHA1 Message Date
Thomas Wolf 3bfeebe722 Less general but avoid hook issues 2020-04-02 10:00:35 +02:00
Thomas Wolf ffed6a8a5d Update modeling_t5.py
Style and quality
2020-04-01 23:28:36 +02:00
Thomas Wolf 426b5106e7 Adding spread_on_devices 2020-04-01 23:21:10 +02:00
+67
View File
@@ -20,6 +20,7 @@ import itertools
import logging
import math
import os
from typing import List, Optional
import torch
import torch.nn.functional as F
@@ -404,6 +405,8 @@ class T5Block(nn.Module):
def __init__(self, config, has_relative_attention_bias=False):
super().__init__()
self.is_decoder = config.is_decoder
self.device = None
self.next_device = None
self.layer = nn.ModuleList()
self.layer.append(T5LayerSelfAttention(config, has_relative_attention_bias=has_relative_attention_bias))
if self.is_decoder:
@@ -422,6 +425,21 @@ class T5Block(nn.Module):
encoder_decoder_position_bias=None,
head_mask=None,
):
if self.device is not None:
(hidden_states,
attention_mask,
position_bias,
encoder_hidden_states,
encoder_attention_mask,
encoder_decoder_position_bias,
head_mask) = tuple(t.to(self.device) for t in (hidden_states,
attention_mask,
position_bias,
encoder_hidden_states,
encoder_attention_mask,
encoder_decoder_position_bias,
head_mask)) # Model parallelism
self_attention_outputs = self.layer[0](
hidden_states, attention_mask=attention_mask, position_bias=position_bias, head_mask=head_mask
)
@@ -445,6 +463,10 @@ class T5Block(nn.Module):
hidden_states = self.layer[2](hidden_states)
outputs = (hidden_states,) + outputs # add attentions if we output them
if self.next_device is not None:
outputs = tuple(t.to(self.device) for t in outputs) # Model parallelism
return outputs # hidden-states, (self-attention weights), (self-attention position bias), (cross-attention weights), (cross-attention position bias)
@@ -541,6 +563,9 @@ class T5Stack(T5PreTrainedModel):
self.init_weights()
def get_block_list(self):
return list(self.block)
def get_input_embeddings(self):
return self.embed_tokens
@@ -773,6 +798,48 @@ class T5Model(T5PreTrainedModel):
self.init_weights()
def spread_on_devices(self, devices: Optional[List] = None):
""" Spread a transformers model on several devices by moving block on several devices (simple model parallelism)
The blocks of the transformers are spread among the given device list
or on all visible CUDA devices if no device list is given.
The first device will host in addition the embeddings and the input/output tensors.
"""
if devices is None and torch.cuda.is_available():
devices = list(range(torch.cuda.device_count()))
if len(devices) < 2:
self.to(devices[0] if devices else None)
return
modules_to_move = set(self.modules)
# Evenly spread the blocks on devices
block_list = self.get_block_list()
group_size = len(block_list) // len(devices)
for i, block in enumerate(block_list):
device = devices[i // group_size]
# Note that we cannot easily use `forward_pre_hook` to move tensors around since this type of hooks currently
# only act on the positional arguments send to the forward pass (PyTorch 1.4.0).
# So you should call your model's forward pass with tensors as positional arguments
# see: https://github.com/pytorch/pytorch/blob/master/torch/nn/modules/module.py#L548-L554
# block.register_forward_pre_hook(lambda module, input: tuple(t.to(device) for t in input))
block.to(device)
block.device = device
modules_to_move.remove(block)
# Take care of brining back the tensors to the first device at the end of the last block's forward
block.next_device = device[0]
# block.register_forward_hook(lambda module, input, output: tuple(t.to(device[0]) for t in output))
# Move the remaining modules (embeddings) on the first device
for module in list(modules_to_move):
module.to(devices[0])
def get_block_list(self):
return list(self.encoder.get_block_list()) + list(self.decoder.get_block_list())
def get_input_embeddings(self):
return self.shared