Packages
Tensor library for Gleam/BEAM with a pure Gleam API, zero-copy views, and optional native acceleration
Current section
Files
Jump to
Current section
Files
src/viva_tensor@generate@speculative.erl
-module(viva_tensor@generate@speculative).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/generate/speculative.gleam").
-export([greedy_sample/1, apply_temperature/2, top_k_filter/2, top_p_filter/2, sample/2, resample_from_residual/2, accept_reject/2, speculative_decode/4, greedy_generate/4]).
-export_type([sampling_config/0, speculative_config/0]).
-if(?OTP_RELEASE >= 27).
-define(MODULEDOC(Str), -moduledoc(Str)).
-define(DOC(Str), -doc(Str)).
-else.
-define(MODULEDOC(Str), -compile([])).
-define(DOC(Str), -compile([])).
-endif.
?MODULEDOC(false).
-type sampling_config() :: {sampling_config, float(), integer(), float()}.
-type speculative_config() :: {speculative_config,
integer(),
sampling_config(),
integer()}.
-file("src/viva_tensor/generate/speculative.gleam", 92).
?DOC(false).
-spec greedy_sample(viva_tensor@tensor:tensor()) -> {ok, integer()} |
{error, viva_tensor@core@error:tensor_error()}.
greedy_sample(Logits) ->
viva_tensor@tensor:try_argmax(Logits).
-file("src/viva_tensor/generate/speculative.gleam", 110).
?DOC(false).
-spec apply_temperature(viva_tensor@tensor:tensor(), float()) -> viva_tensor@tensor:tensor().
apply_temperature(Logits, Temperature) ->
case (Temperature =:= +0.0) orelse (Temperature =:= 1.0) of
true ->
Logits;
false ->
Data = viva_tensor@tensor:to_list(Logits),
Scaled = gleam@list:map(Data, fun(X) -> case Temperature of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X / Gleam@denominator
end end),
{tensor, Scaled, viva_tensor@tensor:shape(Logits)}
end.
-file("src/viva_tensor/generate/speculative.gleam", 615).
?DOC(false).
-spec compare_desc(float(), float()) -> gleam@order:order().
compare_desc(A, B) ->
case A < B of
true ->
gt;
false ->
case A > B of
true ->
lt;
false ->
eq
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 521).
?DOC(false).
-spec ensure_1d(viva_tensor@tensor:tensor()) -> {ok, list(float())} |
{error, viva_tensor@core@error:tensor_error()}.
ensure_1d(T) ->
case viva_tensor@tensor:shape(T) of
[_] ->
{ok, viva_tensor@tensor:to_list(T)};
_ ->
{error,
{dimension_error,
<<"Sampling helpers require a 1-D [vocab_size] tensor"/utf8>>}}
end.
-file("src/viva_tensor/generate/speculative.gleam", 128).
?DOC(false).
-spec top_k_filter(viva_tensor@tensor:tensor(), integer()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
top_k_filter(Logits, K) ->
gleam@result:'try'(
ensure_1d(Logits),
fun(Data) ->
N = erlang:length(Data),
case (K =< 0) orelse (K >= N) of
true ->
{ok, Logits};
false ->
Sorted_desc = gleam@list:sort(
Data,
fun(A, B) -> compare_desc(A, B) end
),
Threshold = begin
_pipe = Sorted_desc,
_pipe@1 = gleam@list:drop(_pipe, K - 1),
_pipe@2 = gleam@list:first(_pipe@1),
gleam@result:unwrap(_pipe@2, +0.0)
end,
Masked = gleam@list:map(
Data,
fun(X) -> case X < Threshold of
true ->
-1.0e30;
false ->
X
end end
),
{ok, {tensor, Masked, viva_tensor@tensor:shape(Logits)}}
end
end
).
-file("src/viva_tensor/generate/speculative.gleam", 597).
?DOC(false).
-spec nucleus_indices(
list({integer(), float()}),
float(),
float(),
list(integer())
) -> list(integer()).
nucleus_indices(Sorted, P, Cum, Acc) ->
case Sorted of
[] ->
Acc;
[{Idx, Prob} | Rest] ->
case Cum >= P of
true ->
Acc;
false ->
nucleus_indices(Rest, P, Cum + Prob, [Idx | Acc])
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 530).
?DOC(false).
-spec softmax_data(list(float())) -> list(float()).
softmax_data(Data) ->
case Data of
[] ->
[];
[First | Rest] ->
Max_val = gleam@list:fold(
Rest,
First,
fun(Acc, X) -> case X > Acc of
true ->
X;
false ->
Acc
end end
),
Shifted = gleam@list:map(
Data,
fun(X@1) -> viva_tensor@core@ffi:exp(X@1 - Max_val) end
),
Sum_exp = gleam@list:fold(
Shifted,
+0.0,
fun(Acc@1, X@2) -> Acc@1 + X@2 end
),
case Sum_exp =< +0.0 of
true ->
gleam@list:map(Data, fun(_) -> +0.0 end);
false ->
gleam@list:map(Shifted, fun(X@3) -> case Sum_exp of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X@3 / Gleam@denominator
end end)
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 171).
?DOC(false).
-spec top_p_filter(viva_tensor@tensor:tensor(), float()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
top_p_filter(Logits, P) ->
gleam@result:'try'(ensure_1d(Logits), fun(Data) -> case P >= 1.0 of
true ->
{ok, Logits};
false ->
Probs = softmax_data(Data),
Indexed = gleam@list:index_map(
Probs,
fun(Prob, I) -> {I, Prob} end
),
Sorted = gleam@list:sort(
Indexed,
fun(A, B) ->
{_, Pa} = A,
{_, Pb} = B,
compare_desc(Pa, Pb)
end
),
Keep_indices = nucleus_indices(Sorted, P, +0.0, []),
Masked = gleam@list:index_map(
Data,
fun(X, I@1) ->
case gleam@list:contains(Keep_indices, I@1) of
true ->
X;
false ->
-1.0e30
end
end
),
{ok, {tensor, Masked, viva_tensor@tensor:shape(Logits)}}
end end).
-file("src/viva_tensor/generate/speculative.gleam", 557).
?DOC(false).
-spec inverse_cdf(list(float()), float(), float(), integer()) -> integer().
inverse_cdf(Probs, U, Cum, Index) ->
case Probs of
[] ->
Index - 1;
[P | Rest] ->
New_cum = Cum + P,
case U =< New_cum of
true ->
Index;
false ->
inverse_cdf(Rest, U, New_cum, Index + 1)
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 570).
?DOC(false).
-spec sample_unit() -> float().
sample_unit() ->
case 2147483648.0 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> erlang:float(gleam@int:random(2147483648)) / Gleam@denominator
end.
-file("src/viva_tensor/generate/speculative.gleam", 552).
?DOC(false).
-spec sample_from_probs(list(float())) -> integer().
sample_from_probs(Probs) ->
U = sample_unit(),
inverse_cdf(Probs, U, +0.0, 0).
-file("src/viva_tensor/generate/speculative.gleam", 207).
?DOC(false).
-spec sample(viva_tensor@tensor:tensor(), sampling_config()) -> {ok, integer()} |
{error, viva_tensor@core@error:tensor_error()}.
sample(Logits, Config) ->
case erlang:element(2, Config) =:= +0.0 of
true ->
greedy_sample(Logits);
false ->
Scaled = apply_temperature(Logits, erlang:element(2, Config)),
gleam@result:'try'(
top_k_filter(Scaled, erlang:element(3, Config)),
fun(After_k) ->
gleam@result:'try'(
top_p_filter(After_k, erlang:element(4, Config)),
fun(After_p) ->
gleam@result:'try'(
ensure_1d(After_p),
fun(Data) -> case Data of
[] ->
{error,
{invalid_shape,
<<"Empty logits tensor"/utf8>>}};
_ ->
{ok,
sample_from_probs(
softmax_data(Data)
)}
end end
)
end
)
end
)
end.
-file("src/viva_tensor/generate/speculative.gleam", 514).
?DOC(false).
-spec bonus_sample(viva_tensor@tensor:tensor(), sampling_config()) -> integer().
bonus_sample(Logits, Sampling) ->
case sample(Logits, Sampling) of
{ok, T} ->
T;
{error, _} ->
0
end.
-file("src/viva_tensor/generate/speculative.gleam", 586).
?DOC(false).
-spec argmax_loop(list(float()), integer(), integer(), float()) -> integer().
argmax_loop(Probs, I, Best_i, Best_v) ->
case Probs of
[] ->
Best_i;
[P | Rest] ->
case P > Best_v of
true ->
argmax_loop(Rest, I + 1, I, P);
false ->
argmax_loop(Rest, I + 1, Best_i, Best_v)
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 582).
?DOC(false).
-spec argmax_of(list(float())) -> integer().
argmax_of(Probs) ->
argmax_loop(Probs, 0, 0, +0.0 - 1.0e30).
-file("src/viva_tensor/generate/speculative.gleam", 432).
?DOC(false).
-spec resample_from_residual(list(float()), list(float())) -> integer().
resample_from_residual(Draft_probs, Target_probs) ->
Residual = begin
_pipe = gleam@list:zip(Target_probs, Draft_probs),
gleam@list:map(
_pipe,
fun(Pair) ->
{T, D} = Pair,
case (T - D) > +0.0 of
true ->
T - D;
false ->
+0.0
end
end
)
end,
Total = gleam@list:fold(Residual, +0.0, fun(Acc, X) -> Acc + X end),
case Total =< +0.0 of
true ->
argmax_of(Target_probs);
false ->
Normalized = gleam@list:map(Residual, fun(X@1) -> case Total of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> X@1 / Gleam@denominator
end end),
sample_from_probs(Normalized)
end.
-file("src/viva_tensor/generate/speculative.gleam", 413).
?DOC(false).
-spec accept_reject(float(), float()) -> boolean().
accept_reject(P_draft, P_target) ->
Ratio = case P_draft =< +0.0 of
true ->
1.0;
false ->
R = case P_draft of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> P_target / Gleam@denominator
end,
case R > 1.0 of
true ->
1.0;
false ->
R
end
end,
U = sample_unit(),
U =< Ratio.
-file("src/viva_tensor/generate/speculative.gleam", 574).
?DOC(false).
-spec list_at(list(float()), integer()) -> float().
list_at(Data, Index) ->
case {Data, Index} of
{[], _} ->
+0.0;
{[X | _], 0} ->
X;
{[_ | Rest], I} ->
list_at(Rest, I - 1)
end.
-file("src/viva_tensor/generate/speculative.gleam", 505).
?DOC(false).
-spec apply_filters(viva_tensor@tensor:tensor(), sampling_config()) -> viva_tensor@tensor:tensor().
apply_filters(Logits, Sampling) ->
Scaled = apply_temperature(Logits, erlang:element(2, Sampling)),
After_k = begin
_pipe = top_k_filter(Scaled, erlang:element(3, Sampling)),
gleam@result:unwrap(_pipe, Scaled)
end,
_pipe@1 = top_p_filter(After_k, erlang:element(4, Sampling)),
gleam@result:unwrap(_pipe@1, After_k).
-file("src/viva_tensor/generate/speculative.gleam", 371).
?DOC(false).
-spec acceptance_loop(
list(integer()),
list(viva_tensor@tensor:tensor()),
list(viva_tensor@tensor:tensor()),
sampling_config(),
integer(),
list(integer())
) -> {list(integer()), gleam@option:option(integer()), integer()}.
acceptance_loop(Drafts, Draft_logits, Verify_logits, Sampling, Index, Accepted) ->
case {Drafts, Draft_logits, Verify_logits} of
{[], _, _} ->
{lists:reverse(Accepted), none, 0};
{_, [], _} ->
{lists:reverse(Accepted), none, 0};
{_, _, []} ->
{lists:reverse(Accepted), none, 0};
{[Tok | D_rest], [D_logits | Dl_rest], [V_logits | Vl_rest]} ->
D_data = viva_tensor@tensor:to_list(
apply_filters(D_logits, Sampling)
),
V_data = viva_tensor@tensor:to_list(
apply_filters(V_logits, Sampling)
),
D_probs = softmax_data(D_data),
V_probs = softmax_data(V_data),
P_draft = list_at(D_probs, Tok),
P_target = list_at(V_probs, Tok),
case accept_reject(P_draft, P_target) of
true ->
acceptance_loop(
D_rest,
Dl_rest,
Vl_rest,
Sampling,
Index + 1,
[Tok | Accepted]
);
false ->
Resampled = resample_from_residual(D_probs, V_probs),
{lists:reverse(Accepted), {some, Index}, Resampled}
end
end.
-file("src/viva_tensor/generate/speculative.gleam", 362).
?DOC(false).
-spec acceptance_phase(
list(integer()),
list(viva_tensor@tensor:tensor()),
list(viva_tensor@tensor:tensor()),
sampling_config()
) -> {list(integer()), gleam@option:option(integer()), integer()}.
acceptance_phase(Drafts, Draft_logits, Verify_logits, Sampling) ->
acceptance_loop(Drafts, Draft_logits, Verify_logits, Sampling, 0, []).
-file("src/viva_tensor/generate/speculative.gleam", 347).
?DOC(false).
-spec verify_phase(
list(integer()),
list(integer()),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
list(viva_tensor@tensor:tensor())
) -> {ok, list(viva_tensor@tensor:tensor())} |
{error, viva_tensor@core@error:tensor_error()}.
verify_phase(Prefix, Drafts, Verify_fn, Acc) ->
case Drafts of
[] ->
{ok, lists:reverse(Acc)};
[D | Rest] ->
gleam@result:'try'(
Verify_fn(Prefix),
fun(Logits) ->
verify_phase(
lists:append(Prefix, [D]),
Rest,
Verify_fn,
[Logits | Acc]
)
end
)
end.
-file("src/viva_tensor/generate/speculative.gleam", 322).
?DOC(false).
-spec draft_phase(
speculative_config(),
list(integer()),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
integer(),
list(integer()),
list(viva_tensor@tensor:tensor())
) -> {ok, {list(integer()), list(viva_tensor@tensor:tensor())}} |
{error, viva_tensor@core@error:tensor_error()}.
draft_phase(Config, Prefix, Draft_fn, Remaining, Acc_tokens, Acc_logits) ->
case Remaining =< 0 of
true ->
{ok, {lists:reverse(Acc_tokens), lists:reverse(Acc_logits)}};
false ->
gleam@result:'try'(
Draft_fn(Prefix),
fun(Logits) ->
gleam@result:'try'(
sample(Logits, erlang:element(3, Config)),
fun(Token) ->
draft_phase(
Config,
lists:append(Prefix, [Token]),
Draft_fn,
Remaining - 1,
[Token | Acc_tokens],
[Logits | Acc_logits]
)
end
)
end
)
end.
-file("src/viva_tensor/generate/speculative.gleam", 267).
?DOC(false).
-spec speculative_loop(
speculative_config(),
list(integer()),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
integer()
) -> {ok, list(integer())} | {error, viva_tensor@core@error:tensor_error()}.
speculative_loop(Config, Tokens, Draft_fn, Verify_fn, Produced) ->
case Produced >= erlang:element(4, Config) of
true ->
{ok, Tokens};
false ->
gleam@result:'try'(
draft_phase(
Config,
Tokens,
Draft_fn,
erlang:element(2, Config),
[],
[]
),
fun(_use0) ->
{Drafts, Draft_logits_list} = _use0,
gleam@result:'try'(
verify_phase(Tokens, Drafts, Verify_fn, []),
fun(Verify_logits_list) ->
{Accepted_tokens, Rejected_at, Residual_token} = acceptance_phase(
Drafts,
Draft_logits_list,
Verify_logits_list,
erlang:element(3, Config)
),
New_tokens = lists:append(Tokens, Accepted_tokens),
Extra_token = case Rejected_at of
{some, _} ->
[Residual_token];
none ->
case gleam@list:last(Verify_logits_list) of
{ok, Last_logits} ->
[bonus_sample(
Last_logits,
erlang:element(3, Config)
)];
{error, _} ->
[]
end
end,
Combined = lists:append(New_tokens, Extra_token),
New_produced = Produced + erlang:length(
Accepted_tokens
),
Final_produced = New_produced + erlang:length(
Extra_token
),
case Final_produced >= erlang:element(4, Config) of
true ->
{ok, Combined};
false ->
speculative_loop(
Config,
Combined,
Draft_fn,
Verify_fn,
Final_produced
)
end
end
)
end
)
end.
-file("src/viva_tensor/generate/speculative.gleam", 258).
?DOC(false).
-spec speculative_decode(
speculative_config(),
list(integer()),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()})
) -> {ok, list(integer())} | {error, viva_tensor@core@error:tensor_error()}.
speculative_decode(Config, Initial_tokens, Draft_fn, Verify_fn) ->
speculative_loop(Config, Initial_tokens, Draft_fn, Verify_fn, 0).
-file("src/viva_tensor/generate/speculative.gleam", 473).
?DOC(false).
-spec greedy_generate_loop(
list(integer()),
integer(),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
gleam@option:option(integer()),
integer()
) -> {ok, list(integer())} | {error, viva_tensor@core@error:tensor_error()}.
greedy_generate_loop(Tokens, Max_new_tokens, Model_fn, Stop_token, Produced) ->
case Produced >= Max_new_tokens of
true ->
{ok, Tokens};
false ->
gleam@result:'try'(
Model_fn(Tokens),
fun(Logits) ->
gleam@result:'try'(
greedy_sample(Logits),
fun(Next) ->
New_tokens = lists:append(Tokens, [Next]),
case Stop_token of
{some, Stop} when Stop =:= Next ->
{ok, New_tokens};
_ ->
greedy_generate_loop(
New_tokens,
Max_new_tokens,
Model_fn,
Stop_token,
Produced + 1
)
end
end
)
end
)
end.
-file("src/viva_tensor/generate/speculative.gleam", 464).
?DOC(false).
-spec greedy_generate(
list(integer()),
integer(),
fun((list(integer())) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}),
gleam@option:option(integer())
) -> {ok, list(integer())} | {error, viva_tensor@core@error:tensor_error()}.
greedy_generate(Initial_tokens, Max_new_tokens, Model_fn, Stop_token) ->
greedy_generate_loop(
Initial_tokens,
Max_new_tokens,
Model_fn,
Stop_token,
0
).