CoolFace
Modelpublic

glaiveai/glaive-function-calling-v2-small

sourceHugging Faceupdated 3y agoView on Hugging Face
15likes69downloads
attention.py278 linesDownload Raw Back to root
1"""Attention layers."""2import math3import warnings4from typing import Optional5import torch6import torch.nn as nn7from einops import rearrange8from torch import nn9from .norm import LPLayerNorm10 11def _reset_is_causal(num_query_tokens: int, num_key_tokens: int, original_is_causal: bool):12    if original_is_causal and num_query_tokens != num_key_tokens:13        if num_query_tokens != 1:14            raise NotImplementedError('MPT does not support query and key with different number of tokens, unless number of query tokens is 1.')15        else:16            return False17    return original_is_causal18 19def scaled_multihead_dot_product_attention(query, key, value, n_heads, softmax_scale=None, attn_bias=None, key_padding_mask=None, is_causal=False, dropout_p=0.0, training=False, needs_weights=False, multiquery=False):20    q = rearrange(query, 'b s (h d) -> b h s d', h=n_heads)21    k = rearrange(key, 'b s (h d) -> b h d s', h=1 if multiquery else n_heads)22    v = rearrange(value, 'b s (h d) -> b h s d', h=1 if multiquery else n_heads)23    min_val = torch.finfo(q.dtype).min24    (b, _, s_q, d) = q.shape25    s_k = k.size(-1)26    if softmax_scale is None:27        softmax_scale = 1 / math.sqrt(d)28    attn_weight = q.matmul(k) * softmax_scale29    if attn_bias is not None:30        if attn_bias.size(-1) != 1 and attn_bias.size(-1) != s_k or (attn_bias.size(-2) != 1 and attn_bias.size(-2) != s_q):31            raise RuntimeError(f'attn_bias (shape: {attn_bias.shape}) is expected to broadcast to shape: {attn_weight.shape}.')32        attn_weight = attn_weight + attn_bias33    if key_padding_mask is not None:34        if attn_bias is not None:35            warnings.warn('Propogating key_padding_mask to the attention module ' + 'and applying it within the attention module can cause ' + 'unneccessary computation/memory usage. Consider integrating ' + 'into attn_bias once and passing that to each attention ' + 'module instead.')36        attn_weight = attn_weight.masked_fill(~key_padding_mask.view((b, 1, 1, s_k)), min_val)37    if is_causal:38        s = max(s_q, s_k)39        causal_mask = attn_weight.new_ones(s, s, dtype=torch.float16)40        causal_mask = causal_mask.tril()41        causal_mask = causal_mask.to(torch.bool)42        causal_mask = ~causal_mask43        causal_mask = causal_mask[-s_q:, -s_k:]44        attn_weight = attn_weight.masked_fill(causal_mask.view(1, 1, s_q, s_k), min_val)45    attn_weight = torch.softmax(attn_weight, dim=-1)46    if dropout_p:47        attn_weight = torch.nn.functional.dropout(attn_weight, p=dropout_p, training=training, inplace=True)48    out = attn_weight.matmul(v)49    out = rearrange(out, 'b h s d -> b s (h d)')50    if needs_weights:51        return (out, attn_weight)52    return (out, None)53 54def check_valid_inputs(*tensors, valid_dtypes=[torch.float16, torch.bfloat16]):55    for tensor in tensors:56        if tensor.dtype not in valid_dtypes:57            raise TypeError(f'tensor.dtype={tensor.dtype!r} must be in valid_dtypes={valid_dtypes!r}.')58        if not tensor.is_cuda:59            raise TypeError(f'Inputs must be cuda tensors (tensor.is_cuda={tensor.is_cuda!r}).')60 61def flash_attn_fn(query, key, value, n_heads, softmax_scale=None, attn_bias=None, key_padding_mask=None, is_causal=False, dropout_p=0.0, training=False, needs_weights=False, multiquery=False):62    try:63        from flash_attn import bert_padding, flash_attn_interface64    except:65        raise RuntimeError('Please install flash-attn==1.0.3.post0')66    check_valid_inputs(query, key, value)67    if attn_bias is not None:68        raise NotImplementedError(f'attn_bias not implemented for flash attn.')69    (batch_size, seqlen) = query.shape[:2]70    if key_padding_mask is None:71        key_padding_mask = torch.ones_like(key[:, :, 0], dtype=torch.bool)72    query_padding_mask = key_padding_mask[:, -query.size(1):]73    (query_unpad, indices_q, cu_seqlens_q, max_seqlen_q) = bert_padding.unpad_input(query, query_padding_mask)74    query_unpad = rearrange(query_unpad, 'nnz (h d) -> nnz h d', h=n_heads)75    (key_unpad, _, cu_seqlens_k, max_seqlen_k) = bert_padding.unpad_input(key, key_padding_mask)76    key_unpad = rearrange(key_unpad, 'nnz (h d) -> nnz h d', h=1 if multiquery else n_heads)77    (value_unpad, _, _, _) = bert_padding.unpad_input(value, key_padding_mask)78    value_unpad = rearrange(value_unpad, 'nnz (h d) -> nnz h d', h=1 if multiquery else n_heads)79    if multiquery:80        key_unpad = key_unpad.expand(key_unpad.size(0), n_heads, key_unpad.size(-1))81        value_unpad = value_unpad.expand(value_unpad.size(0), n_heads, value_unpad.size(-1))82    dropout_p = dropout_p if training else 0.083    reset_is_causal = _reset_is_causal(query.size(1), key.size(1), is_causal)84    output_unpad = flash_attn_interface.flash_attn_unpadded_func(query_unpad, key_unpad, value_unpad, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p, softmax_scale=softmax_scale, causal=reset_is_causal, return_attn_probs=needs_weights)85    output = bert_padding.pad_input(rearrange(output_unpad, 'nnz h d -> nnz (h d)'), indices_q, batch_size, seqlen)86    return (output, None)87 88def triton_flash_attn_fn(query, key, value, n_heads, softmax_scale=None, attn_bias=None, key_padding_mask=None, is_causal=False, dropout_p=0.0, training=False, needs_weights=False, multiquery=False):89    try:90        from flash_attn import flash_attn_triton91    except:92        raise RuntimeError('Please install flash-attn==1.0.3.post0 and triton==2.0.0.dev20221202')93    check_valid_inputs(query, key, value)94    if dropout_p:95        raise NotImplementedError(f'Dropout not implemented for attn_impl: triton.')96    if needs_weights:97        raise NotImplementedError(f'attn_impl: triton cannot return attn weights.')98    if key_padding_mask is not None:99        warnings.warn('Propagating key_padding_mask to the attention module ' + 'and applying it within the attention module can cause ' + 'unnecessary computation/memory usage. Consider integrating ' + 'into attn_bias once and passing that to each attention ' + 'module instead.')100        (b_size, s_k) = key_padding_mask.shape[:2]101        if attn_bias is None:102            attn_bias = query.new_zeros(b_size, 1, 1, s_k)103        attn_bias = attn_bias.masked_fill(~key_padding_mask.view((b_size, 1, 1, s_k)), torch.finfo(query.dtype).min)104    query = rearrange(query, 'b s (h d) -> b s h d', h=n_heads)105    key = rearrange(key, 'b s (h d) -> b s h d', h=1 if multiquery else n_heads)106    value = rearrange(value, 'b s (h d) -> b s h d', h=1 if multiquery else n_heads)107    if multiquery:108        key = key.expand(*key.shape[:2], n_heads, key.size(-1))109        value = value.expand(*value.shape[:2], n_heads, value.size(-1))110    reset_is_causal = _reset_is_causal(query.size(1), key.size(1), is_causal)111    attn_output = flash_attn_triton.flash_attn_func(query, key, value, attn_bias, reset_is_causal, softmax_scale)112    output = attn_output.view(*attn_output.shape[:2], -1)113    return (output, None)114 115class MultiheadAttention(nn.Module):116    """Multi-head self attention.117 118    Using torch or triton attention implemetation enables user to also use119    additive bias.120    """121 122    def __init__(self, d_model: int, n_heads: int, attn_impl: str='triton', clip_qkv: Optional[float]=None, qk_ln: bool=False, softmax_scale: Optional[float]=None, attn_pdrop: float=0.0, low_precision_layernorm: bool=False, verbose: int=0, device: Optional[str]=None):123        super().__init__()124        self.attn_impl = attn_impl125        self.clip_qkv = clip_qkv126        self.qk_ln = qk_ln127        self.d_model = d_model128        self.n_heads = n_heads129        self.softmax_scale = softmax_scale130        if self.softmax_scale is None:131            self.softmax_scale = 1 / math.sqrt(self.d_model / self.n_heads)132        self.attn_dropout_p = attn_pdrop133        self.Wqkv = nn.Linear(self.d_model, 3 * self.d_model, device=device)134        fuse_splits = (d_model, 2 * d_model)135        self.Wqkv._fused = (0, fuse_splits)136        if self.qk_ln:137            layernorm_class = LPLayerNorm if low_precision_layernorm else nn.LayerNorm138            self.q_ln = layernorm_class(self.d_model, device=device)139            self.k_ln = layernorm_class(self.d_model, device=device)140        if self.attn_impl == 'flash':141            self.attn_fn = flash_attn_fn142        elif self.attn_impl == 'triton':143            self.attn_fn = triton_flash_attn_fn144            if verbose:145                warnings.warn('While `attn_impl: triton` can be faster than `attn_impl: flash` ' + 'it uses more memory. When training larger models this can trigger ' + 'alloc retries which hurts performance. If encountered, we recommend ' + 'using `attn_impl: flash` if your model does not use `alibi` or `prefix_lm`.')146        elif self.attn_impl == 'torch':147            self.attn_fn = scaled_multihead_dot_product_attention148            if torch.cuda.is_available() and verbose:149                warnings.warn('Using `attn_impl: torch`. If your model does not use `alibi` or ' + '`prefix_lm` we recommend using `attn_impl: flash` otherwise ' + 'we recommend using `attn_impl: triton`.')150        else:151            raise ValueError(f'attn_impl={attn_impl!r} is an invalid setting.')152        self.out_proj = nn.Linear(self.d_model, self.d_model, device=device)153        self.out_proj._is_residual = True154 155    def forward(self, x, past_key_value=None, attn_bias=None, attention_mask=None, is_causal=True, needs_weights=False):156        qkv = self.Wqkv(x)157        if self.clip_qkv:158            qkv.clamp_(min=-self.clip_qkv, max=self.clip_qkv)159        (query, key, value) = qkv.chunk(3, dim=2)160        key_padding_mask = attention_mask161        if self.qk_ln:162            dtype = query.dtype163            query = self.q_ln(query).to(dtype)164            key = self.k_ln(key).to(dtype)165        if past_key_value is not None:166            if len(past_key_value) != 0:167                key = torch.cat([past_key_value[0], key], dim=1)168                value = torch.cat([past_key_value[1], value], dim=1)169            past_key_value = (key, value)170        if attn_bias is not None:171            attn_bias = attn_bias[:, :, -query.size(1):, -key.size(1):]172        (context, attn_weights) = self.attn_fn(query, key, value, self.n_heads, softmax_scale=self.softmax_scale, attn_bias=attn_bias, key_padding_mask=key_padding_mask, is_causal=is_causal, dropout_p=self.attn_dropout_p, training=self.training, needs_weights=needs_weights)173        return (self.out_proj(context), attn_weights, past_key_value)174 175class MultiQueryAttention(nn.Module):176    """Multi-Query self attention.177 178    Using torch or triton attention implemetation enables user to also use179    additive bias.180    """181 182    def __init__(self, d_model: int, n_heads: int, attn_impl: str='triton', clip_qkv: Optional[float]=None, qk_ln: bool=False, softmax_scale: Optional[float]=None, attn_pdrop: float=0.0, low_precision_layernorm: bool=False, verbose: int=0, device: Optional[str]=None):183        super().__init__()184        self.attn_impl = attn_impl185        self.clip_qkv = clip_qkv186        self.qk_ln = qk_ln187        self.d_model = d_model188        self.n_heads = n_heads189        self.head_dim = d_model // n_heads190        self.softmax_scale = softmax_scale191        if self.softmax_scale is None:192            self.softmax_scale = 1 / math.sqrt(self.head_dim)193        self.attn_dropout_p = attn_pdrop194        self.Wqkv = nn.Linear(d_model, d_model + 2 * self.head_dim, device=device)195        fuse_splits = (d_model, d_model + self.head_dim)196        self.Wqkv._fused = (0, fuse_splits)197        if self.qk_ln:198            layernorm_class = LPLayerNorm if low_precision_layernorm else nn.LayerNorm199            self.q_ln = layernorm_class(d_model, device=device)200            self.k_ln = layernorm_class(self.head_dim, device=device)201        if self.attn_impl == 'flash':202            self.attn_fn = flash_attn_fn203        elif self.attn_impl == 'triton':204            self.attn_fn = triton_flash_attn_fn205            if verbose:206                warnings.warn('While `attn_impl: triton` can be faster than `attn_impl: flash` ' + 'it uses more memory. When training larger models this can trigger ' + 'alloc retries which hurts performance. If encountered, we recommend ' + 'using `attn_impl: flash` if your model does not use `alibi` or `prefix_lm`.')207        elif self.attn_impl == 'torch':208            self.attn_fn = scaled_multihead_dot_product_attention209            if torch.cuda.is_available() and verbose:210                warnings.warn('Using `attn_impl: torch`. If your model does not use `alibi` or ' + '`prefix_lm` we recommend using `attn_impl: flash` otherwise ' + 'we recommend using `attn_impl: triton`.')211        else:212            raise ValueError(f'attn_impl={attn_impl!r} is an invalid setting.')213        self.out_proj = nn.Linear(self.d_model, self.d_model, device=device)214        self.out_proj._is_residual = True215 216    def forward(self, x, past_key_value=None, attn_bias=None, attention_mask=None, is_causal=True, needs_weights=False):217        qkv = self.Wqkv(x)218        if self.clip_qkv:219            qkv.clamp_(min=-self.clip_qkv, max=self.clip_qkv)220        (query, key, value) = qkv.split([self.d_model, self.head_dim, self.head_dim], dim=2)221        key_padding_mask = attention_mask222        if self.qk_ln:223            dtype = query.dtype224            query = self.q_ln(query).to(dtype)225            key = self.k_ln(key).to(dtype)226        if past_key_value is not None:227            if len(past_key_value) != 0:228                key = torch.cat([past_key_value[0], key], dim=1)229                value = torch.cat([past_key_value[1], value], dim=1)230            past_key_value = (key, value)231        if attn_bias is not None:232            attn_bias = attn_bias[:, :, -query.size(1):, -key.size(1):]233        (context, attn_weights) = self.attn_fn(query, key, value, self.n_heads, softmax_scale=self.softmax_scale, attn_bias=attn_bias, key_padding_mask=key_padding_mask, is_causal=is_causal, dropout_p=self.attn_dropout_p, training=self.training, needs_weights=needs_weights, multiquery=True)234        return (self.out_proj(context), attn_weights, past_key_value)235 236def attn_bias_shape(attn_impl, n_heads, seq_len, alibi, prefix_lm, causal, use_sequence_id):237    if attn_impl == 'flash':238        return None239    elif attn_impl in ['torch', 'triton']:240        if alibi:241            if (prefix_lm or not causal) or use_sequence_id:242                return (1, n_heads, seq_len, seq_len)243            return (1, n_heads, 1, seq_len)244        elif prefix_lm or use_sequence_id:245            return (1, 1, seq_len, seq_len)246        return None247    else:248        raise ValueError(f'attn_impl={attn_impl!r} is an invalid setting.')249 250def build_attn_bias(attn_impl, attn_bias, n_heads, seq_len, causal=False, alibi=False, alibi_bias_max=8):251    if attn_impl == 'flash':252        return None253    elif attn_impl in ['torch', 'triton']:254        if alibi:255            (device, dtype) = (attn_bias.device, attn_bias.dtype)256            attn_bias = attn_bias.add(build_alibi_bias(n_heads, seq_len, full=not causal, alibi_bias_max=alibi_bias_max, device=device, dtype=dtype))257        return attn_bias258    else:259        raise ValueError(f'attn_impl={attn_impl!r} is an invalid setting.')260 261def gen_slopes(n_heads, alibi_bias_max=8, device=None):262    _n_heads = 2 ** math.ceil(math.log2(n_heads))263    m = torch.arange(1, _n_heads + 1, dtype=torch.float32, device=device)264    m = m.mul(alibi_bias_max / _n_heads)265    slopes = 1.0 / torch.pow(2, m)266    if _n_heads != n_heads:267        slopes = torch.concat([slopes[1::2], slopes[::2]])[:n_heads]268    return slopes.view(1, n_heads, 1, 1)269 270def build_alibi_bias(n_heads, seq_len, full=False, alibi_bias_max=8, device=None, dtype=None):271    alibi_bias = torch.arange(1 - seq_len, 1, dtype=torch.int32, device=device).view(1, 1, 1, seq_len)272    if full:273        alibi_bias = alibi_bias - torch.arange(1 - seq_len, 1, dtype=torch.int32, device=device).view(1, 1, seq_len, 1)274        alibi_bias = alibi_bias.abs().mul(-1)275    slopes = gen_slopes(n_heads, alibi_bias_max, device=device)276    alibi_bias = alibi_bias * slopes277    return alibi_bias.to(dtype=dtype)278ATTN_CLASS_REGISTRY = {'multihead_attention': MultiheadAttention, 'multiquery_attention': MultiQueryAttention}