-module(viva_tensor@nn@pool). -compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]). -define(FILEPATH, "src/viva_tensor/nn/pool.gleam"). -export([dropout_init/1, dropout_forward/3, max_pool_1d_forward/2, avg_pool_1d_forward/2, adaptive_avg_pool_2d_forward/2, adaptive_avg_pool_1d_forward/2, upsample_forward/2]). -export_type([dropout/0, max_pool1d_config/0, avg_pool1d_config/0, adaptive_avg_pool2d_config/0, adaptive_avg_pool1d_config/0, upsample_mode/0, upsample_config/0, acc/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 dropout() :: {dropout, float()}. -type max_pool1d_config() :: {max_pool1d_config, integer(), integer(), integer()}. -type avg_pool1d_config() :: {avg_pool1d_config, integer(), integer(), integer()}. -type adaptive_avg_pool2d_config() :: {adaptive_avg_pool2d_config, integer(), integer()}. -type adaptive_avg_pool1d_config() :: {adaptive_avg_pool1d_config, integer()}. -type upsample_mode() :: nearest | bilinear. -type upsample_config() :: {upsample_config, integer(), upsample_mode()}. -type acc() :: acc_none | {acc_some, float()}. -file("src/viva_tensor/nn/pool.gleam", 94). ?DOC(false). -spec dropout_init(float()) -> dropout(). dropout_init(P) -> {dropout, P}. -file("src/viva_tensor/nn/pool.gleam", 114). ?DOC(false). -spec dropout_forward(dropout(), viva_tensor@tensor:tensor(), boolean()) -> viva_tensor@tensor:tensor(). dropout_forward(Layer, Input, Training) -> case Training of false -> Input; true -> case erlang:element(2, Layer) of P when P =< +0.0 -> Input; P@1 when P@1 >= 1.0 -> Shp = viva_tensor@tensor:shape(Input), N = erlang:length(viva_tensor@tensor:to_list(Input)), Zeros = gleam@list:repeat(+0.0, N), case viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Zeros), Shp ) of {ok, T} -> T; {error, _} -> Input end; P@2 -> Keep_prob = 1.0 - P@2, Scale = case Keep_prob of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> 1.0 / Gleam@denominator end, Resolution = 1000000, Threshold = erlang:round( Keep_prob * erlang:float(Resolution) ), Shp@1 = viva_tensor@tensor:shape(Input), Masked = gleam@list:map( viva_tensor@tensor:to_list(Input), fun(V) -> case gleam@int:random(Resolution) < Threshold of true -> V * Scale; false -> +0.0 end end ), case viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Masked), Shp@1 ) of {ok, T@1} -> T@1; {error, _} -> Input end end end. -file("src/viva_tensor/nn/pool.gleam", 730). ?DOC(false). -spec join_strings(list(binary()), binary()) -> binary(). join_strings(Parts, Sep) -> case Parts of [] -> <<""/utf8>>; [S] -> S; [S@1 | Rest] -> <<<>/binary, (join_strings(Rest, Sep))/binary>> end. -file("src/viva_tensor/nn/pool.gleam", 726). ?DOC(false). -spec shape_to_string(list(integer())) -> binary(). shape_to_string(Shape) -> <<<<"["/utf8, (join_strings( gleam@list:map(Shape, fun erlang:integer_to_binary/1), <<", "/utf8>> ))/binary>>/binary, "]"/utf8>>. -file("src/viva_tensor/nn/pool.gleam", 340). ?DOC(false). -spec sum_window( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), float() ) -> float(). sum_window(Arr, Base, Start, K, Kernel_size, Length, Acc) -> case K >= Kernel_size of true -> Acc; false -> Idx = Start + K, V = case (Idx >= 0) andalso (Idx < Length) of true -> viva_tensor@core@ffi:array_get(Arr, Base + Idx); false -> +0.0 end, sum_window(Arr, Base, Start, K + 1, Kernel_size, Length, Acc + V) end. -file("src/viva_tensor/nn/pool.gleam", 306). ?DOC(false). -spec max_window( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), acc() ) -> float(). max_window(Arr, Base, Start, K, Kernel_size, Length, Acc) -> case K >= Kernel_size of true -> case Acc of acc_none -> +0.0; {acc_some, V} -> V end; false -> Idx = Start + K, V@1 = case (Idx >= 0) andalso (Idx < Length) of true -> viva_tensor@core@ffi:array_get(Arr, Base + Idx); false -> +0.0 end, Next = case Acc of acc_none -> {acc_some, V@1}; {acc_some, M} -> case V@1 > M of true -> {acc_some, V@1}; false -> {acc_some, M} end end, max_window(Arr, Base, Start, K + 1, Kernel_size, Length, Next) end. -file("src/viva_tensor/nn/pool.gleam", 283). ?DOC(false). -spec pool_1d_window( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), boolean(), float() ) -> float(). pool_1d_window(Arr, Base, Start, Kernel_size, Length, Is_max, Kf) -> case Is_max of true -> max_window(Arr, Base, Start, 0, Kernel_size, Length, acc_none); false -> Sum = sum_window(Arr, Base, Start, 0, Kernel_size, Length, +0.0), case Kf of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> Sum / Gleam@denominator end end. -file("src/viva_tensor/nn/pool.gleam", 742). ?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/nn/pool.gleam", 738). ?DOC(false). -spec range_int(integer(), integer()) -> list(integer()). range_int(From, To) -> range_loop(From, To, []). -file("src/viva_tensor/nn/pool.gleam", 206). ?DOC(false). -spec pool_1d_forward( viva_tensor@tensor:tensor(), integer(), integer(), integer(), boolean() ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. pool_1d_forward(Input, Kernel_size, Stride, Padding, Is_max) -> Shp = viva_tensor@tensor:shape(Input), case Shp of [Batch, Channels, Length] -> case (Stride =< 0) orelse (Kernel_size =< 0) of true -> {error, {invalid_shape, <<<<<<"pool_1d: kernel_size and stride must be positive, got kernel="/utf8, (erlang:integer_to_binary(Kernel_size))/binary>>/binary, " stride="/utf8>>/binary, (erlang:integer_to_binary(Stride))/binary>>}}; false -> Padded_length = Length + (2 * Padding), Out_length = (case Stride of 0 -> 0; Gleam@denominator -> (Padded_length - Kernel_size) div Gleam@denominator end) + 1, case Out_length =< 0 of true -> {error, {invalid_shape, <<<<<<<<<<<<<<<<<<"pool_1d: invalid output length "/utf8, (erlang:integer_to_binary( Out_length ))/binary>>/binary, " for input length "/utf8>>/binary, (erlang:integer_to_binary( Length ))/binary>>/binary, ", kernel "/utf8>>/binary, (erlang:integer_to_binary( Kernel_size ))/binary>>/binary, ", stride "/utf8>>/binary, (erlang:integer_to_binary( Stride ))/binary>>/binary, ", padding "/utf8>>/binary, (erlang:integer_to_binary(Padding))/binary>>}}; false -> Arr = viva_tensor@core@ffi:list_to_array( viva_tensor@tensor:to_list(Input) ), Kf = erlang:float(Kernel_size), Out = begin _pipe = range_int(0, Batch - 1), gleam@list:flat_map( _pipe, fun(B) -> _pipe@1 = range_int(0, Channels - 1), gleam@list:flat_map( _pipe@1, fun(C) -> Base = ((B * Channels) * Length) + (C * Length), _pipe@2 = range_int( 0, Out_length - 1 ), gleam@list:map( _pipe@2, fun(O) -> Start = (O * Stride) - Padding, pool_1d_window( Arr, Base, Start, Kernel_size, Length, Is_max, Kf ) end ) end ) end ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, Channels, Out_length] ) end end; _ -> {error, {invalid_shape, <<"pool_1d: input must have shape [batch, channels, length], got "/utf8, (shape_to_string(Shp))/binary>>}} end. -file("src/viva_tensor/nn/pool.gleam", 168). ?DOC(false). -spec max_pool_1d_forward(max_pool1d_config(), viva_tensor@tensor:tensor()) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. max_pool_1d_forward(Config, Input) -> pool_1d_forward( Input, erlang:element(2, Config), erlang:element(3, Config), erlang:element(4, Config), true ). -file("src/viva_tensor/nn/pool.gleam", 193). ?DOC(false). -spec avg_pool_1d_forward(avg_pool1d_config(), viva_tensor@tensor:tensor()) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. avg_pool_1d_forward(Config, Input) -> pool_1d_forward( Input, erlang:element(2, Config), erlang:element(3, Config), erlang:element(4, Config), false ). -file("src/viva_tensor/nn/pool.gleam", 465). ?DOC(false). -spec sum_2d_cols( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), float() ) -> float(). sum_2d_cols(Arr, Base, W_in, H, W, W_end, Acc) -> case W >= W_end of true -> Acc; false -> V = viva_tensor@core@ffi:array_get(Arr, (Base + (H * W_in)) + W), sum_2d_cols(Arr, Base, W_in, H, W + 1, W_end, Acc + V) end. -file("src/viva_tensor/nn/pool.gleam", 446). ?DOC(false). -spec sum_2d_rows( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), integer(), float() ) -> float(). sum_2d_rows(Arr, Base, W_in, H, H_end, W_start, W_end, Acc) -> case H >= H_end of true -> Acc; false -> Row_sum = sum_2d_cols(Arr, Base, W_in, H, W_start, W_end, +0.0), sum_2d_rows( Arr, Base, W_in, H + 1, H_end, W_start, W_end, Acc + Row_sum ) end. -file("src/viva_tensor/nn/pool.gleam", 434). ?DOC(false). -spec sum_2d_window( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), integer() ) -> float(). sum_2d_window(Arr, Base, W_in, H_start, H_end, W_start, W_end) -> sum_2d_rows(Arr, Base, W_in, H_start, H_end, W_start, W_end, +0.0). -file("src/viva_tensor/nn/pool.gleam", 559). ?DOC(false). -spec ceil_div(integer(), integer()) -> integer(). ceil_div(A, B) -> case B of 0 -> 0; Gleam@denominator -> ((A + B) - 1) div Gleam@denominator end. -file("src/viva_tensor/nn/pool.gleam", 553). ?DOC(false). -spec adaptive_range(integer(), integer(), integer()) -> {integer(), integer()}. adaptive_range(I, In_size, Out_size) -> Start = case Out_size of 0 -> 0; Gleam@denominator -> I * In_size div Gleam@denominator end, End = ceil_div((I + 1) * In_size, Out_size), {Start, End}. -file("src/viva_tensor/nn/pool.gleam", 375). ?DOC(false). -spec adaptive_avg_pool_2d_forward( adaptive_avg_pool2d_config(), viva_tensor@tensor:tensor() ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. adaptive_avg_pool_2d_forward(Config, Input) -> Shp = viva_tensor@tensor:shape(Input), case Shp of [Batch, Channels, H_in, W_in] -> case (erlang:element(2, Config) =< 0) orelse (erlang:element( 3, Config ) =< 0) of true -> {error, {invalid_shape, <<<<<<"adaptive_avg_pool_2d: output dims must be positive, got "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>/binary, "x"/utf8>>/binary, (erlang:integer_to_binary( erlang:element(3, Config) ))/binary>>}}; false -> Arr = viva_tensor@core@ffi:list_to_array( viva_tensor@tensor:to_list(Input) ), Out_h = erlang:element(2, Config), Out_w = erlang:element(3, Config), Out = begin _pipe = range_int(0, Batch - 1), gleam@list:flat_map( _pipe, fun(B) -> _pipe@1 = range_int(0, Channels - 1), gleam@list:flat_map( _pipe@1, fun(C) -> Base = (((B * Channels) * H_in) * W_in) + ((C * H_in) * W_in), _pipe@2 = range_int(0, Out_h - 1), gleam@list:flat_map( _pipe@2, fun(Oh) -> {H_start, H_end} = adaptive_range( Oh, H_in, Out_h ), _pipe@3 = range_int( 0, Out_w - 1 ), gleam@list:map( _pipe@3, fun(Ow) -> {W_start, W_end} = adaptive_range( Ow, W_in, Out_w ), Count = (H_end - H_start) * (W_end - W_start), Sum = sum_2d_window( Arr, Base, W_in, H_start, H_end, W_start, W_end ), case erlang:float(Count) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> Sum / Gleam@denominator end end ) end ) end ) end ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, Channels, Out_h, Out_w] ) end; _ -> {error, {invalid_shape, <<"adaptive_avg_pool_2d: input must have shape [batch, channels, H, W], got "/utf8, (shape_to_string(Shp))/binary>>}} end. -file("src/viva_tensor/nn/pool.gleam", 536). ?DOC(false). -spec sum_1d_range( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), float() ) -> float(). sum_1d_range(Arr, Base, I, End, Acc) -> case I >= End of true -> Acc; false -> sum_1d_range( Arr, Base, I + 1, End, Acc + viva_tensor@core@ffi:array_get(Arr, Base + I) ) end. -file("src/viva_tensor/nn/pool.gleam", 493). ?DOC(false). -spec adaptive_avg_pool_1d_forward( adaptive_avg_pool1d_config(), viva_tensor@tensor:tensor() ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. adaptive_avg_pool_1d_forward(Config, Input) -> Shp = viva_tensor@tensor:shape(Input), case Shp of [Batch, Channels, Length] -> case erlang:element(2, Config) =< 0 of true -> {error, {invalid_shape, <<"adaptive_avg_pool_1d: output_size must be positive, got "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>}}; false -> Arr = viva_tensor@core@ffi:list_to_array( viva_tensor@tensor:to_list(Input) ), Out_size = erlang:element(2, Config), Out = begin _pipe = range_int(0, Batch - 1), gleam@list:flat_map( _pipe, fun(B) -> _pipe@1 = range_int(0, Channels - 1), gleam@list:flat_map( _pipe@1, fun(C) -> Base = ((B * Channels) * Length) + (C * Length), _pipe@2 = range_int(0, Out_size - 1), gleam@list:map( _pipe@2, fun(O) -> {Start, End} = adaptive_range( O, Length, Out_size ), Count = End - Start, Sum = sum_1d_range( Arr, Base, Start, End, +0.0 ), case erlang:float(Count) of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> Sum / Gleam@denominator end end ) end ) end ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, Channels, Out_size] ) end; _ -> {error, {invalid_shape, <<"adaptive_avg_pool_1d: input must have shape [batch, channels, length], got "/utf8, (shape_to_string(Shp))/binary>>}} end. -file("src/viva_tensor/nn/pool.gleam", 704). ?DOC(false). -spec interp_neighbors(float(), integer()) -> {integer(), integer(), float()}. interp_neighbors(Src, Size) -> Clamped = case Src < +0.0 of true -> +0.0; false -> case Src > erlang:float(Size - 1) of true -> erlang:float(Size - 1); false -> Src end end, Lo = erlang:trunc(Clamped), Hi = case (Lo + 1) >= Size of true -> Size - 1; false -> Lo + 1 end, Frac = Clamped - erlang:float(Lo), {Lo, Hi, Frac}. -file("src/viva_tensor/nn/pool.gleam", 663). ?DOC(false). -spec upsample_bilinear_compute( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), integer(), integer() ) -> list(float()). upsample_bilinear_compute(Arr, Batch, Channels, H_in, W_in, H_out, W_out, Scale) -> Scale_f = erlang:float(Scale), _pipe = range_int(0, Batch - 1), gleam@list:flat_map(_pipe, fun(B) -> _pipe@1 = range_int(0, Channels - 1), gleam@list:flat_map( _pipe@1, fun(C) -> Base = (((B * Channels) * H_in) * W_in) + ((C * H_in) * W_in), _pipe@2 = range_int(0, H_out - 1), gleam@list:flat_map( _pipe@2, fun(Oh) -> Src_h = (case Scale_f of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator -> (erlang:float(Oh) + 0.5) / Gleam@denominator end) - 0.5, {H0, H1, Dh} = interp_neighbors(Src_h, H_in), _pipe@3 = range_int(0, W_out - 1), gleam@list:map( _pipe@3, fun(Ow) -> Src_w = (case Scale_f of +0.0 -> +0.0; -0.0 -> -0.0; Gleam@denominator@1 -> (erlang:float(Ow) + 0.5) / Gleam@denominator@1 end) - 0.5, {W0, W1, Dw} = interp_neighbors(Src_w, W_in), V00 = viva_tensor@core@ffi:array_get( Arr, (Base + (H0 * W_in)) + W0 ), V01 = viva_tensor@core@ffi:array_get( Arr, (Base + (H0 * W_in)) + W1 ), V10 = viva_tensor@core@ffi:array_get( Arr, (Base + (H1 * W_in)) + W0 ), V11 = viva_tensor@core@ffi:array_get( Arr, (Base + (H1 * W_in)) + W1 ), Top = (V00 * (1.0 - Dw)) + (V01 * Dw), Bot = (V10 * (1.0 - Dw)) + (V11 * Dw), (Top * (1.0 - Dh)) + (Bot * Dh) end ) end ) end ) end). -file("src/viva_tensor/nn/pool.gleam", 635). ?DOC(false). -spec upsample_nearest_compute( viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), integer(), integer() ) -> list(float()). upsample_nearest_compute(Arr, Batch, Channels, H_in, W_in, H_out, W_out, Scale) -> _pipe = range_int(0, Batch - 1), gleam@list:flat_map(_pipe, fun(B) -> _pipe@1 = range_int(0, Channels - 1), gleam@list:flat_map( _pipe@1, fun(C) -> Base = (((B * Channels) * H_in) * W_in) + ((C * H_in) * W_in), _pipe@2 = range_int(0, H_out - 1), gleam@list:flat_map( _pipe@2, fun(Oh) -> Ih = case Scale of 0 -> 0; Gleam@denominator -> Oh div Gleam@denominator end, _pipe@3 = range_int(0, W_out - 1), gleam@list:map( _pipe@3, fun(Ow) -> Iw = case Scale of 0 -> 0; Gleam@denominator@1 -> Ow div Gleam@denominator@1 end, viva_tensor@core@ffi:array_get( Arr, (Base + (Ih * W_in)) + Iw ) end ) end ) end ) end). -file("src/viva_tensor/nn/pool.gleam", 577). ?DOC(false). -spec upsample_forward(upsample_config(), viva_tensor@tensor:tensor()) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. upsample_forward(Config, Input) -> Shp = viva_tensor@tensor:shape(Input), case Shp of [Batch, Channels, H_in, W_in] -> case erlang:element(2, Config) =< 0 of true -> {error, {invalid_shape, <<"upsample: scale_factor must be positive, got "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>}}; false -> H_out = H_in * erlang:element(2, Config), W_out = W_in * erlang:element(2, Config), Arr = viva_tensor@core@ffi:list_to_array( viva_tensor@tensor:to_list(Input) ), Out = case erlang:element(3, Config) of nearest -> upsample_nearest_compute( Arr, Batch, Channels, H_in, W_in, H_out, W_out, erlang:element(2, Config) ); bilinear -> upsample_bilinear_compute( Arr, Batch, Channels, H_in, W_in, H_out, W_out, erlang:element(2, Config) ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, Channels, H_out, W_out] ) end; _ -> {error, {invalid_shape, <<"upsample: input must have shape [batch, channels, H, W], got "/utf8, (shape_to_string(Shp))/binary>>}} end.