bigscience/petals-api
18
1"""2Utility operations used in the the BLOOM model3Based on https://github.com/huggingface/transformers/commit/ca2a55e9dfb245527b5e1c954fec6ffbb7aef07b4See commit history for authorship.5"""6import math7 8import torch9import torch.autograd10import torch.nn.functional as F11from torch import nn12 13 14def split_tensor_along_last_dim(tensor, num_partitions, contiguous_split_chunks=False):15 """Split a tensor along its last dimension.16 17 Args:18 tensor: ([`torch.tensor`], *required*):19 input tensor to split20 num_partitions ([`int`], *required*):21 number of partitions to split the tensor22 contiguous_split_chunks ([`bool`], *optional*, default=`False`)::23 If True, make each chunk contiguous in memory.24 """25 # Get the size and dimension.26 last_dim = tensor.dim() - 127 numerator, denominator = tensor.size()[last_dim], num_partitions28 if not (numerator % denominator == 0):29 raise ValueError(f"{numerator} is not divisible by {denominator}")30 last_dim_size = numerator // denominator31 # Split.32 tensor_list = torch.split(tensor, last_dim_size, dim=last_dim)33 # Note: torch.split does not create contiguous tensors by default.34 if contiguous_split_chunks:35 return tuple(chunk.contiguous() for chunk in tensor_list)36 37 return tensor_list38 39 40def attention_mask_func(attention_scores, attention_mask, causal_mask):41 if attention_mask.dtype == torch.bool:42 attention_mask_bool = ~attention_mask43 else:44 attention_mask_bool = (1 - attention_mask).bool()45 46 query_length, key_length, n_heads = attention_scores.size(2), attention_scores.size(3), attention_scores.size(1)47 padded_causal_mask = (48 attention_mask_bool[:, None, key_length - query_length : key_length, None]49 + ~causal_mask[:, :, key_length - query_length : key_length, :key_length]50 ).bool()51 padded_causal_mask = padded_causal_mask + attention_mask_bool[:, None, None, :key_length].bool()52 # Make use of floats53 return (54 attention_scores.masked_fill_(padded_causal_mask.expand(-1, n_heads, -1, -1), -10000.0),55 padded_causal_mask,56 )57 58 59def build_alibi_tensor(60 max_seq_len: int, n_head: int, dtype: torch.dtype = torch.bfloat16, device: torch.device = torch.device("cpu")61) -> torch.Tensor:62 """63 Link to paper: https://arxiv.org/abs/2108.12409 Alibi tensor is not causal as the original paper mentions, it64 relies on a translation invariance of softmax for quick implementation: with l being a tensor, and a fixed value65 `softmax(l+a) = softmax(l)`. Based on66 https://github.com/ofirpress/attention_with_linear_biases/blob/a35aaca144e0eb6b789dfcb46784c4b8e31b7983/fairseq/models/transformer.py#L74267 Args:68 Returns tensor shaped (n_head, 1, max_seq_len)69 max_seq_len: (`int`, *required*):70 max sequence length71 n_head: (`int`, *required*):72 number of heads73 dtype: (`torch.dtype`, *optional*, default=`torch.bfloat16`):74 dtype of the output tensor75 device: (`torch.device`, *optional*, default=`torch.device('cpu')`):76 device of the output alibi tensor77 """78 closest_power_of_2 = 2 ** math.floor(math.log2(n_head))79 base = torch.tensor(2 ** (-(2 ** -(math.log2(closest_power_of_2) - 3))), device=device, dtype=torch.float32)80 powers = torch.arange(1, 1 + closest_power_of_2, device=device, dtype=torch.int32)81 slopes = torch.pow(base, powers)82 83 if closest_power_of_2 != n_head:84 extra_base = torch.tensor(85 2 ** (-(2 ** -(math.log2(2 * closest_power_of_2) - 3))), device=device, dtype=torch.float3286 )87 num_remaining_heads = min(closest_power_of_2, n_head - closest_power_of_2)88 extra_powers = torch.arange(1, 1 + 2 * num_remaining_heads, 2, device=device, dtype=torch.int32)89 slopes = torch.cat([slopes, torch.pow(extra_base, extra_powers)], dim=0)90 91 lengths = torch.arange(max_seq_len, device=device, dtype=torch.int32)92 return (slopes.view(-1, 1, 1) * lengths.view(1, 1, -1)).to(dtype)93 94 95def pre_process_alibi_for_pad(alibi: torch.Tensor, attention_mask: torch.Tensor):96 """97 Args:98 Pre-process the alibi tensor for padding.99 alibi: ([`torch.tensor`], *required*):100 alibi tensor to pre-process101 attention_mask: ([`torch.tensor`], *required*):102 attention mask to pre-process103 """104 assert attention_mask.shape.ndim == 2, "mask should be [batch_size, seq_length]"105 unpadded_indices = torch.relu(attention_mask.cumsum(dim=1) - 1)106 # ^-- [batch, max_len], values correspond to element indices after removing padding107 # We shift the alibi tensor + replace all the values where attention_mask==0.0 by 0108 alibi = alibi.take_along_dim(unpadded_indices.unsqueeze(0), -1) * attention_mask.unsqueeze(0)109 return alibi.reshape(alibi.shape[0] * alibi.shape[1], 1, -1)110 111 112def dropout_add(x, residual, prob, training):113 """114 Dropout add function115 116 Args:117 x (`torch.tensor`, *required*):118 input tensor119 residual (`torch.tensor`, *rquired*):120 esidual tensor121 prob (`float`, *required*):122 dropout probability123 training (`bool`, *required*):124 training mode125 """126 out = nn.functional.dropout(x, p=prob, training=training)127 out = residual + out128 return out129 130 131def bloom_gelu_forward(x):132 """133 Custom bias GELU function. Adapted from Megatron-DeepSpeed code. Here we use a simple implementation (inference) to134 make the model jitable.135 136 Args:137 x (`torch.tensor`, *required*):138 input hidden states139 """140 return x * 0.5 * (1.0 + torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x)))141 142 143def bloom_gelu_back(g, x):144 """145 gradient of tanh approximation of gelu gradient of actual gelu is: 0.5 * (1. + torch.erf(x * 0.70710678)) +146 0.3989423 * x * torch.exp(-0.5 * x * x)147 148 Args:149 g (`torch.tensor`, *required*):150 gradient output tensor151 x (`torch.tensor`, *required*):152 input tensor153 """154 x = x[0] # x is a tuple of 1 element, needs to unpack it first155 tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))156 # sqrt(2/pi) * 3 * 0.044715 -> 0.1070322243157 ff = 0.5 * x * ((1 - tanh_out * tanh_out) * (0.79788456 + 0.1070322243 * x * x)) + 0.5 * (1 + tanh_out)158 return ff * g159 160 161class GeLUFunction(torch.autograd.Function):162 @staticmethod163 def forward(ctx, input):164 ctx.save_for_backward(input)165 return bloom_gelu_forward(input)166 167 @staticmethod168 def backward(ctx, grad_output):169 input = ctx.saved_tensors170 tmp = bloom_gelu_back(grad_output, input)171 return tmp172 173 174class BloomGelu(nn.Module):175 """176 BloomBiasGelu wrapper function that make use of the simple function on inference mode to make the model177 torchscriptable and use the autograd function in training mode to get the accurate results of the gradients Partly178 copied from Megatron-DeepSpeed code and adapted for our needs179 180 See here why autograd functions are not torchscriptable: https://github.com/pytorch/pytorch/issues/22329181 182 """183 184 def __init__(self):185 super().__init__()186 187 def forward(self, x):188 if self.training:189 return GeLUFunction.apply(x)190 else:191 return bloom_gelu_forward(x)192 193 194class BloomScaledSoftmax(nn.Module):195 """196 fused operation: scaling + mask + softmax197 198 Args:199 input_in_fp16 (`bool`, *required*):200 flag to indicate if input in fp16 data format.201 input_in_bf16 (`bool`, *required*):202 flag to indicate if input in bf16 data format.203 scaled_masked_softmax_fusion (`bool`, *required*):204 flag to indicate user want to use softmax fusion205 mask_func (`function`, *required*):206 mask function to be applied.207 softmax_in_fp32 (`bool`, *required*):208 if true, softmax in performed at fp32 precision.209 scale (`float`, *required*):210 scaling factor used in input tensor scaling.211 """212 213 def __init__(self, scaled_masked_softmax_fusion, mask_func, softmax_in_fp32, scale):214 super().__init__()215 self.scaled_masked_softmax_fusion = scaled_masked_softmax_fusion216 self.mask_func = mask_func217 self.softmax_in_fp32 = softmax_in_fp32218 self.scale = scale219 220 if not (self.scale is None or softmax_in_fp32):221 raise ValueError("softmax should be in fp32 when scaled")222 223 def forward(self, input, mask, max_positions):224 input_dtype = input.dtype225 input_in_16bit = input_dtype in [torch.float16, torch.bfloat16]226 softmax_dtype = torch.float32 if self.softmax_in_fp32 else input_dtype227 228 if self.scale is not None:229 input = input * self.scale230 231 if mask is None:232 mask = torch.ones(input.shape[0], max_positions, dtype=torch.bool, device=input.device)233 234 mask = mask.to(input.device)235 causal_mask = (236 torch.tril(torch.ones((max_positions, max_positions), dtype=torch.bool))237 .view(1, 1, max_positions, max_positions)238 .to(input.device)239 )240 mask_output, padded_causal_mask = self.mask_func(input, mask, causal_mask)241 probs = F.softmax(mask_output, dim=-1, dtype=softmax_dtype) * (~padded_causal_mask)242 243 if input_in_16bit and self.softmax_in_fp32:244 probs = probs.to(dtype=input_dtype)245 246 return probs247 