Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3bfeebe722 | ||
|
|
ffed6a8a5d | ||
|
|
426b5106e7 |
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user