-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, <<<<<<<<<>/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.