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@diffusion@samplers.erl
-module(viva_tensor@diffusion@samplers).
-compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]).
-define(FILEPATH, "src/viva_tensor/diffusion/samplers.gleam").
-export([build_schedule/1, ddpm_step/4, ddim_step/5, sample/4]).
-export_type([noise_schedule/0, sampler_config/0, scheduler_state/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 noise_schedule() :: {linear_schedule, float(), float(), integer()} |
{cosine_schedule, integer()}.
-type sampler_config() :: {sampler_config, noise_schedule(), float()}.
-type scheduler_state() :: {scheduler_state,
list(float()),
list(float()),
list(float()),
integer()}.
-file("src/viva_tensor/diffusion/samplers.gleam", 142).
?DOC(false).
-spec betas_from_alpha_bars(list(float()), float()) -> list(float()).
betas_from_alpha_bars(Alpha_bars, Prev_seed) ->
_pipe = erlang:element(
1,
gleam@list:fold(
Alpha_bars,
{[], Prev_seed},
fun(Acc, Ab) ->
{Out, Prev} = Acc,
Beta = case Prev =< +0.0 of
true ->
0.999;
false ->
gleam@float:clamp(1.0 - (case Prev of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Ab / Gleam@denominator
end), +0.0, 0.999)
end,
{[Beta | Out], Ab}
end
)
),
lists:reverse(_pipe).
-file("src/viva_tensor/diffusion/samplers.gleam", 420).
?DOC(false).
-spec range_loop(integer(), integer(), list(integer())) -> list(integer()).
range_loop(From, To, Acc) ->
case From > To of
true ->
lists:reverse(Acc);
false ->
range_loop(From + 1, To, [From | Acc])
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 416).
?DOC(false).
-spec range_int(integer(), integer()) -> list(integer()).
range_int(From, To) ->
range_loop(From, To, []).
-file("src/viva_tensor/diffusion/samplers.gleam", 122).
?DOC(false).
-spec cumulative_product(list(float())) -> list(float()).
cumulative_product(Xs) ->
_pipe = erlang:element(
1,
gleam@list:fold(
Xs,
{[], 1.0},
fun(Acc, Value) ->
{Accum, Running} = Acc,
Next = Running * Value,
{[Next | Accum], Next}
end
)
),
lists:reverse(_pipe).
-file("src/viva_tensor/diffusion/samplers.gleam", 111).
?DOC(false).
-spec schedule_from_betas(list(float()), integer()) -> scheduler_state().
schedule_from_betas(Betas, Num_steps) ->
Alphas = gleam@list:map(Betas, fun(B) -> 1.0 - B end),
Alpha_bars = cumulative_product(Alphas),
{scheduler_state, Betas, Alphas, Alpha_bars, Num_steps}.
-file("src/viva_tensor/diffusion/samplers.gleam", 131).
?DOC(false).
-spec linspace_floats(float(), float(), integer()) -> list(float()).
linspace_floats(Start, Stop, Steps) ->
case Steps =< 1 of
true ->
[Start];
false ->
Delta = case erlang:float(Steps - 1) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> (Stop - Start) / Gleam@denominator
end,
_pipe = range_int(0, Steps - 1),
gleam@list:map(
_pipe,
fun(I) -> Start + (Delta * erlang:float(I)) end
)
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 71).
?DOC(false).
-spec build_schedule(noise_schedule()) -> scheduler_state().
build_schedule(Schedule) ->
case Schedule of
{linear_schedule, Beta_start, Beta_end, Num_steps} ->
Betas = linspace_floats(Beta_start, Beta_end, Num_steps),
schedule_from_betas(Betas, Num_steps);
{cosine_schedule, Num_steps@1} ->
S_offset = 0.008,
Denom = 1.0 + S_offset,
F = fun(T) ->
X = ((case Denom of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> ((case erlang:float(Num_steps@1) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> erlang:float(T) / Gleam@denominator
end) + S_offset) / Gleam@denominator@1
end) * 3.14159265358979323846) / 2.0,
C = viva_tensor@core@ffi:cos(X),
C * C
end,
F0 = F(0),
Alpha_bars = begin
_pipe = range_int(1, Num_steps@1),
gleam@list:map(_pipe, fun(T@1) -> case F0 =< +0.0 of
true ->
1.0;
false ->
case F0 of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@2 -> F(T@1) / Gleam@denominator@2
end
end end)
end,
Betas@1 = betas_from_alpha_bars(Alpha_bars, 1.0),
Alphas = gleam@list:map(Betas@1, fun(B) -> 1.0 - B end),
{scheduler_state, Betas@1, Alphas, Alpha_bars, Num_steps@1}
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 410).
?DOC(false).
-spec standard_normal() -> float().
standard_normal() ->
U1 = gleam@float:max(viva_tensor@core@ffi:random_uniform(), 1.0e-12),
U2 = viva_tensor@core@ffi:random_uniform(),
viva_tensor@core@ffi:sqrt(-2.0 * viva_tensor@core@ffi:log(U1)) * viva_tensor@core@ffi:cos(
(2.0 * 3.14159265358979323846) * U2
).
-file("src/viva_tensor/diffusion/samplers.gleam", 380).
?DOC(false).
-spec ensure_same_length(binary(), list(float()), list(float())) -> {ok, nil} |
{error, viva_tensor@core@error:tensor_error()}.
ensure_same_length(Op, A, B) ->
case erlang:length(A) =:= erlang:length(B) of
true ->
{ok, nil};
false ->
{error,
{invalid_shape,
<<<<<<<<<<Op/binary,
": x_t and model_pred have different element counts ("/utf8>>/binary,
(erlang:integer_to_binary(erlang:length(A)))/binary>>/binary,
" vs "/utf8>>/binary,
(erlang:integer_to_binary(erlang:length(B)))/binary>>/binary,
")"/utf8>>}}
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 399).
?DOC(false).
-spec at(list(float()), integer(), float()) -> float().
at(Xs, I, Default) ->
case gleam@list:drop(Xs, I) of
[V | _] ->
V;
[] ->
Default
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 366).
?DOC(false).
-spec validate_step(scheduler_state(), integer()) -> {ok, nil} |
{error, viva_tensor@core@error:tensor_error()}.
validate_step(State, T) ->
case (T < 0) orelse (T >= erlang:element(5, State)) of
true ->
{error,
{dimension_error,
<<<<<<<<"diffusion step: t="/utf8,
(erlang:integer_to_binary(T))/binary>>/binary,
" out of range for "/utf8>>/binary,
(erlang:integer_to_binary(erlang:element(5, State)))/binary>>/binary,
"-step schedule"/utf8>>}};
false ->
{ok, nil}
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 172).
?DOC(false).
-spec ddpm_step(
scheduler_state(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
integer()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
ddpm_step(State, X_t, Model_pred, T) ->
gleam@result:'try'(
validate_step(State, T),
fun(_) ->
Alpha_t = at(erlang:element(3, State), T, 1.0),
Alpha_bar_t = at(erlang:element(4, State), T, 1.0),
Beta_t = at(erlang:element(2, State), T, +0.0),
Alpha_bar_prev = case T =:= 0 of
true ->
1.0;
false ->
at(erlang:element(4, State), T - 1, 1.0)
end,
Sqrt_alpha_t = viva_tensor@core@ffi:sqrt(
gleam@float:max(Alpha_t, +0.0)
),
Sqrt_one_minus_bar = viva_tensor@core@ffi:sqrt(
gleam@float:max(1.0 - Alpha_bar_t, +0.0)
),
Xs = viva_tensor@tensor:to_list(X_t),
Preds = viva_tensor@tensor:to_list(Model_pred),
gleam@result:'try'(
ensure_same_length(<<"ddpm_step"/utf8>>, Xs, Preds),
fun(_) ->
Coef = case Sqrt_one_minus_bar =< +0.0 of
true ->
+0.0;
false ->
case Sqrt_one_minus_bar of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> (1.0 - Alpha_t) / Gleam@denominator
end
end,
Inv_sqrt_alpha = case Sqrt_alpha_t =< +0.0 of
true ->
+0.0;
false ->
case Sqrt_alpha_t of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> 1.0 / Gleam@denominator@1
end
end,
Raw_var = case (1.0 - Alpha_bar_t) =< +0.0 of
true ->
+0.0;
false ->
case (1.0 - Alpha_bar_t) of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@2 -> Beta_t * (1.0 - Alpha_bar_prev)
/ Gleam@denominator@2
end
end,
Variance = gleam@float:max(Raw_var, +0.0),
Sigma = case T =:= 0 of
true ->
+0.0;
false ->
viva_tensor@core@ffi:sqrt(Variance)
end,
Next = gleam@list:map(
gleam@list:zip(Xs, Preds),
fun(Pair) ->
{X, Eps} = Pair,
Mean = Inv_sqrt_alpha * (X - (Coef * Eps)),
Z = case Sigma =< +0.0 of
true ->
+0.0;
false ->
standard_normal()
end,
Mean + (Sigma * Z)
end
),
{ok, {tensor, Next, viva_tensor@tensor:shape(X_t)}}
end
)
end
).
-file("src/viva_tensor/diffusion/samplers.gleam", 243).
?DOC(false).
-spec ddim_step(
scheduler_state(),
viva_tensor@tensor:tensor(),
viva_tensor@tensor:tensor(),
integer(),
float()
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
ddim_step(State, X_t, Model_pred, T, Eta) ->
gleam@result:'try'(
validate_step(State, T),
fun(_) ->
Alpha_bar_t = at(erlang:element(4, State), T, 1.0),
Alpha_bar_prev = case T =:= 0 of
true ->
1.0;
false ->
at(erlang:element(4, State), T - 1, 1.0)
end,
One_minus_bar = gleam@float:max(1.0 - Alpha_bar_t, +0.0),
One_minus_prev = gleam@float:max(1.0 - Alpha_bar_prev, +0.0),
Ratio = case Alpha_bar_prev =< +0.0 of
true ->
+0.0;
false ->
gleam@float:max(1.0 - (case Alpha_bar_prev of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator -> Alpha_bar_t / Gleam@denominator
end), +0.0)
end,
Sigma_sq = case One_minus_bar =< +0.0 of
true ->
+0.0;
false ->
(case One_minus_bar of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@1 -> (Eta * Eta) * One_minus_prev / Gleam@denominator@1
end) * Ratio
end,
Sigma_sq@1 = gleam@float:max(Sigma_sq, +0.0),
Sigma = case T =:= 0 of
true ->
+0.0;
false ->
viva_tensor@core@ffi:sqrt(Sigma_sq@1)
end,
Dir_coef = viva_tensor@core@ffi:sqrt(
gleam@float:max(One_minus_prev - Sigma_sq@1, +0.0)
),
Sqrt_alpha_bar = viva_tensor@core@ffi:sqrt(
gleam@float:max(Alpha_bar_t, +0.0)
),
Sqrt_alpha_bar_prev = viva_tensor@core@ffi:sqrt(
gleam@float:max(Alpha_bar_prev, +0.0)
),
Sqrt_one_minus_bar = viva_tensor@core@ffi:sqrt(One_minus_bar),
Xs = viva_tensor@tensor:to_list(X_t),
Preds = viva_tensor@tensor:to_list(Model_pred),
gleam@result:'try'(
ensure_same_length(<<"ddim_step"/utf8>>, Xs, Preds),
fun(_) ->
Next = gleam@list:map(
gleam@list:zip(Xs, Preds),
fun(Pair) ->
{X, Eps} = Pair,
Pred_x0 = case Sqrt_alpha_bar =< +0.0 of
true ->
+0.0;
false ->
case Sqrt_alpha_bar of
+0.0 -> +0.0;
-0.0 -> -0.0;
Gleam@denominator@2 -> (X - (Sqrt_one_minus_bar
* Eps))
/ Gleam@denominator@2
end
end,
Dir = Dir_coef * Eps,
Z = case Sigma =< +0.0 of
true ->
+0.0;
false ->
standard_normal()
end,
((Sqrt_alpha_bar_prev * Pred_x0) + Dir) + (Sigma * Z)
end
),
{ok, {tensor, Next, viva_tensor@tensor:shape(X_t)}}
end
)
end
).
-file("src/viva_tensor/diffusion/samplers.gleam", 338).
?DOC(false).
-spec sampling_loop(
sampler_config(),
scheduler_state(),
viva_tensor@tensor:tensor(),
integer(),
fun((viva_tensor@tensor:tensor(), integer()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()})
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
sampling_loop(Config, State, X_t, T, Model_fn) ->
case T < 0 of
true ->
{ok, X_t};
false ->
gleam@result:'try'(
Model_fn(X_t, T),
fun(Pred) ->
gleam@result:'try'(case erlang:element(3, Config) =< +0.0 of
true ->
ddim_step(State, X_t, Pred, T, +0.0);
false ->
case erlang:element(3, Config) >= 1.0 of
true ->
ddpm_step(State, X_t, Pred, T);
false ->
ddim_step(
State,
X_t,
Pred,
T,
erlang:element(3, Config)
)
end
end, fun(Next) ->
sampling_loop(Config, State, Next, T - 1, Model_fn)
end)
end
)
end.
-file("src/viva_tensor/diffusion/samplers.gleam", 406).
?DOC(false).
-spec element_count(list(integer())) -> integer().
element_count(Shape) ->
gleam@list:fold(Shape, 1, fun(Acc, Dim) -> Acc * Dim end).
-file("src/viva_tensor/diffusion/samplers.gleam", 309).
?DOC(false).
-spec sample(
sampler_config(),
scheduler_state(),
list(integer()),
fun((viva_tensor@tensor:tensor(), integer()) -> {ok,
viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()})
) -> {ok, viva_tensor@tensor:tensor()} |
{error, viva_tensor@core@error:tensor_error()}.
sample(Config, State, Shape, Model_fn) ->
case erlang:element(5, State) =< 0 of
true ->
{error,
{invalid_shape,
<<<<"sample: scheduler has "/utf8,
(erlang:integer_to_binary(erlang:element(5, State)))/binary>>/binary,
" steps"/utf8>>}};
false ->
Total = element_count(Shape),
case Total =< 0 of
true ->
{error,
{invalid_shape, <<"sample: empty target shape"/utf8>>}};
false ->
X_t = {tensor,
begin
_pipe = range_int(1, Total),
gleam@list:map(
_pipe,
fun(_) -> standard_normal() end
)
end,
Shape},
sampling_loop(
Config,
State,
X_t,
erlang:element(5, State) - 1,
Model_fn
)
end
end.