Team Ai
Apppublic

bigscience/petals-api

sourceHugging Faceupdated 4y agoView on Hugging Face
18likes
ops.py247 linesDownload Raw Back to bloom
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