Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
239 changes: 231 additions & 8 deletions lib/bumblebee/layers.ex
Original file line number Diff line number Diff line change
Expand Up @@ -1239,18 +1239,206 @@ defmodule Bumblebee.Layers do
def rotary_embedding(query, key, position_ids, attention_mask, size, opts \\ []) do
opts = Keyword.validate!(opts, [:name, :scaling_strategy, max_positions: 2048, base: 10_000])

scaling_strategy =
normalize_rotary_embedding_scaling_strategy(opts[:scaling_strategy], size, opts)

output =
Axon.layer(
&apply_rotary_embedding/5,
[query, key, position_ids, Axon.optional(attention_mask)],
[size: size] ++ opts
[size: size, scaling_strategy: scaling_strategy] ++
Keyword.drop(opts, [:scaling_strategy])
)

unwrap_tuple(output, 2)
end

defp normalize_rotary_embedding_scaling_strategy(nil, _size, _opts), do: nil

defp normalize_rotary_embedding_scaling_strategy(
%{type: :yarn} = strategy,
size,
opts
) do
required_keys = [
:factor,
:original_max_positions,
:beta_fast,
:beta_slow
]

case Enum.find(required_keys, &(not Map.has_key?(strategy, &1))) do
nil ->
validate_factor!(strategy.factor, :yarn)

strategy =
strategy
|> Map.put_new_lazy(:attention_factor, fn -> 0.1 * :math.log(strategy.factor) + 1.0 end)
|> Map.put_new(:truncate, true)

validate_yarn_strategy!(strategy)

low =
yarn_correction_dim(
strategy.beta_fast,
size,
opts[:base],
strategy.original_max_positions
)

high =
yarn_correction_dim(
strategy.beta_slow,
size,
opts[:base],
strategy.original_max_positions
)

{low, high} =
if strategy.truncate do
{:math.floor(low), :math.ceil(high)}
else
{low, high}
end

low = max(low, 0)
high = min(high, size - 1)
high = if low == high, do: high + 0.001, else: high

Map.merge(strategy, %{
correction_low: low,
correction_high: high
})

key ->
raise ArgumentError, "YaRN rotary embedding scaling strategy is missing #{inspect(key)}"
end
end

defp normalize_rotary_embedding_scaling_strategy(
%{type: :longrope, short_factor: short_factor, long_factor: long_factor} = strategy,
size,
_opts
) do
expected_size = div(size, 2)

if longrope_factors?(short_factor, expected_size) and
longrope_factors?(long_factor, expected_size) do
strategy
else
raise ArgumentError,
"LongRoPE factors must each have #{expected_size} entries for a rotary size of #{size}"
end
end

defp normalize_rotary_embedding_scaling_strategy(
%{type: :dynamic, factor: factor} = strategy,
size,
_opts
) do
validate_factor!(factor, :dynamic)

if size <= 2 do
raise ArgumentError, "dynamic RoPE requires a rotary size greater than 2"
end

strategy
end

defp normalize_rotary_embedding_scaling_strategy(
%{type: :linear, factor: factor} = strategy,
_size,
_opts
) do
validate_factor!(factor, :linear)
strategy
end

defp normalize_rotary_embedding_scaling_strategy(%{type: type}, _size, _opts)
when type in [:linear, :dynamic] do
raise ArgumentError, "#{type} RoPE requires :factor to be a number >= 1"
end

defp normalize_rotary_embedding_scaling_strategy(
%{
type: :llama3,
factor: factor,
low_frequency_factor: low,
high_frequency_factor: high,
original_max_positions: original_max_positions
} = strategy,
_size,
_opts
) do
validate_factor!(factor, :llama3)
validate_positive_number!(low, :low_frequency_factor, :llama3)
validate_positive_number!(high, :high_frequency_factor, :llama3)
validate_positive_integer!(original_max_positions, :original_max_positions, :llama3)

if high <= low do
raise ArgumentError, "llama3 RoPE requires :high_frequency_factor > :low_frequency_factor"
end

strategy
end

defp normalize_rotary_embedding_scaling_strategy(strategy, _size, _opts) do
raise ArgumentError, "unsupported rotary embedding scaling strategy: #{inspect(strategy)}"
end

defp yarn_correction_dim(num_rotations, size, base, original_max_positions) do
size * :math.log(original_max_positions / (num_rotations * 2 * :math.pi())) /
(2 * :math.log(base))
end

defp validate_yarn_strategy!(strategy) do
validate_factor!(strategy.factor, :yarn)
validate_positive_integer!(strategy.original_max_positions, :original_max_positions, :yarn)
validate_positive_number!(strategy.beta_fast, :beta_fast, :yarn)
validate_positive_number!(strategy.beta_slow, :beta_slow, :yarn)
validate_positive_number!(strategy.attention_factor, :attention_factor, :yarn)

unless is_boolean(strategy.truncate) do
raise ArgumentError, "YaRN RoPE requires :truncate to be a boolean"
end

if strategy.beta_fast < strategy.beta_slow do
raise ArgumentError, "YaRN RoPE requires :beta_fast >= :beta_slow"
end
end

defp longrope_factors?(factors, expected_size) do
is_list(factors) and length(factors) == expected_size and
Enum.all?(factors, &(is_number(&1) and &1 > 0))
end

defp validate_factor!(factor, type) do
if is_number(factor) and factor >= 1.0 do
:ok
else
raise ArgumentError, "#{type} RoPE requires :factor to be a number >= 1"
end
end

defp validate_positive_number!(value, key, type) do
if is_number(value) and value > 0 do
:ok
else
raise ArgumentError, "#{type} RoPE requires #{inspect(key)} to be a positive number"
end
end

defp validate_positive_integer!(value, key, type) do
if is_integer(value) and value > 0 do
:ok
else
raise ArgumentError, "#{type} RoPE requires #{inspect(key)} to be a positive integer"
end
end

deftransformp create_sinusoidal_positions(
sequence_length,
context_length,
max_positions,
size,
base,
Expand All @@ -1266,15 +1454,48 @@ defmodule Bumblebee.Layers do
position = Nx.divide(position, factor)
positions_cos_sin(position, inv_frequency)

%{type: :dynamic, factor: factor} when sequence_length > max_positions ->
base =
base
|> Nx.multiply(factor * sequence_length / max_positions - (factor - 1))
|> Nx.pow(size / (size - 2))
%{type: :dynamic, factor: factor} ->
context_length = Nx.max(context_length, max_positions)

scale =
context_length
|> Nx.multiply(factor)
|> Nx.divide(max_positions)
|> Nx.subtract(factor - 1)

base = Nx.multiply(base, Nx.pow(scale, size / (size - 2)))

inv_frequency = inv_frequency(base, range)
positions_cos_sin(position, inv_frequency)

%{
type: :yarn,
factor: factor,
attention_factor: attention_factor,
correction_low: low,
correction_high: high
} ->
inv_frequency_extrapolation = inv_frequency(base, range)
inv_frequency_interpolation = Nx.divide(inv_frequency_extrapolation, factor)

ramp =
Nx.iota({div(size, 2)})
|> Nx.subtract(low)
|> Nx.divide(high - low)
|> Nx.clip(0.0, 1.0)

inv_frequency_extrapolation_factor = Nx.subtract(1.0, ramp)

inv_frequency =
Nx.multiply(
inv_frequency_interpolation,
Nx.subtract(1.0, inv_frequency_extrapolation_factor)
)
|> Nx.add(Nx.multiply(inv_frequency_extrapolation, inv_frequency_extrapolation_factor))

{cos, sin} = positions_cos_sin(position, inv_frequency)
{Nx.multiply(cos, attention_factor), Nx.multiply(sin, attention_factor)}

%{
type: :longrope,
short_factor: short_factor,
Expand Down Expand Up @@ -1394,17 +1615,19 @@ defmodule Bumblebee.Layers do
_other -> Nx.axis_size(attention_mask, 1)
end

position_ids = Nx.as_type(position_ids, :s64)
context_length = Nx.add(Nx.reduce_max(position_ids), 1)

{cos, sin} =
create_sinusoidal_positions(
sequence_length,
context_length,
opts[:max_positions],
opts[:size],
opts[:base],
opts[:scaling_strategy]
)

position_ids = Nx.as_type(position_ids, :s64)

cos = cos |> Nx.take(position_ids) |> Nx.new_axis(2) |> Nx.as_type(Nx.type(query))
sin = sin |> Nx.take(position_ids) |> Nx.new_axis(2) |> Nx.as_type(Nx.type(query))

Expand Down
18 changes: 18 additions & 0 deletions lib/bumblebee/layers/decoder.ex
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,24 @@ defmodule Bumblebee.Layers.Decoder do

alias Bumblebee.Layers

@doc """
Rejects cached inference with bidirectional self-attention, since adding
tokens changes the hidden states of previously processed tokens.
"""
def validate_attention_cache(cache, false), do: cache

def validate_attention_cache(cache, true) do
Axon.layer(
fn cache, _opts ->
case cache do
%Axon.None{} -> cache
_ -> raise ArgumentError, "decoding cache is not supported with bidirectional attention"
end
end,
[Axon.optional(cache)]
)
end

defmodule Cache do
@moduledoc false

Expand Down
Loading
Loading