(n_heads, alibi_bias_max=8)
| 20 | |
| 21 | |
| 22 | def gen_slopes(n_heads, alibi_bias_max=8): |
| 23 | _n_heads = 2 ** math.ceil(math.log2(n_heads)) |
| 24 | m = torch.arange(1, _n_heads + 1, dtype=torch.float32) |
| 25 | m = m.mul(alibi_bias_max / _n_heads) |
| 26 | slopes = 1.0 / torch.pow(2, m) |
| 27 | if _n_heads != n_heads: |
| 28 | slopes = torch.concat([slopes[1::2], slopes[::2]])[:n_heads] |
| 29 | return slopes.view(1, n_heads, 1, 1) |
| 30 | |
| 31 | |
| 32 | def build_alibi_bias( |