-module(viva_tensor@nn@conv). -compile([no_auto_import, nowarn_unused_vars, nowarn_unused_function, nowarn_nomatch, inline]). -define(FILEPATH, "src/viva_tensor/nn/conv.gleam"). -export([conv1d_init/5, conv1d_forward/2, conv3d_init/5, conv3d_forward/2, conv_transpose_2d_init/6, conv_transpose_2d_forward/2]). -export_type([conv1d_config/0, conv3d_config/0, conv_transpose2d_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 conv1d_config() :: {conv1d_config, integer(), integer(), integer(), integer(), integer(), viva_tensor@tensor:tensor(), gleam@option:option(viva_tensor@tensor:tensor())}. -type conv3d_config() :: {conv3d_config, integer(), integer(), {integer(), integer(), integer()}, {integer(), integer(), integer()}, {integer(), integer(), integer()}, viva_tensor@tensor:tensor(), gleam@option:option(viva_tensor@tensor:tensor())}. -type conv_transpose2d_config() :: {conv_transpose2d_config, integer(), integer(), {integer(), integer()}, {integer(), integer()}, {integer(), integer()}, {integer(), integer()}, viva_tensor@tensor:tensor(), gleam@option:option(viva_tensor@tensor:tensor())}. -file("src/viva_tensor/nn/conv.gleam", 93). ?DOC(false). -spec conv1d_init(integer(), integer(), integer(), integer(), integer()) -> conv1d_config(). conv1d_init(In_channels, Out_channels, Kernel_size, Stride, Padding) -> {conv1d_config, In_channels, Out_channels, Kernel_size, Stride, Padding, viva_tensor@tensor:zeros([Out_channels, In_channels, Kernel_size]), {some, viva_tensor@tensor:zeros([Out_channels])}}. -file("src/viva_tensor/nn/conv.gleam", 769). ?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/conv.gleam", 765). ?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/conv.gleam", 753). ?DOC(false). -spec sum_range_acc(integer(), integer(), fun((integer()) -> float()), float()) -> float(). sum_range_acc(Start, End, F, Acc) -> case Start > End of true -> Acc; false -> sum_range_acc(Start + 1, End, F, Acc + F(Start)) end. -file("src/viva_tensor/nn/conv.gleam", 749). ?DOC(false). -spec sum_over_range(integer(), integer(), fun((integer()) -> float())) -> float(). sum_over_range(Start, End, F) -> sum_range_acc(Start, End, F, +0.0). -file("src/viva_tensor/nn/conv.gleam", 781). ?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/conv.gleam", 777). ?DOC(false). -spec range_int(integer(), integer()) -> list(integer()). range_int(From, To) -> range_loop(From, To, []). -file("src/viva_tensor/nn/conv.gleam", 222). ?DOC(false). -spec conv1d_compute( list(float()), list(float()), list(float()), integer(), integer(), integer(), integer(), integer(), integer(), integer() ) -> list(float()). conv1d_compute( In_data, W_data, Bias_data, Batch, In_c, Padded_length, Out_c, Kernel_size, Stride, Out_length ) -> In_arr = viva_tensor@core@ffi:list_to_array(In_data), W_arr = viva_tensor@core@ffi:list_to_array(W_data), Bias_arr = viva_tensor@core@ffi:list_to_array(Bias_data), _pipe = range_int(0, Batch - 1), gleam@list:flat_map(_pipe, fun(B) -> _pipe@1 = range_int(0, Out_c - 1), gleam@list:flat_map( _pipe@1, fun(Oc) -> Bias_v = viva_tensor@core@ffi:array_get(Bias_arr, Oc), _pipe@2 = range_int(0, Out_length - 1), gleam@list:map( _pipe@2, fun(O) -> Start = O * Stride, Sum = sum_over_range( 0, In_c - 1, fun(Ic) -> sum_over_range( 0, Kernel_size - 1, fun(K) -> In_idx = ((((B * In_c) * Padded_length) + (Ic * Padded_length)) + Start) + K, W_idx = (((Oc * In_c) * Kernel_size) + (Ic * Kernel_size)) + K, viva_tensor@core@ffi:array_get( In_arr, In_idx ) * viva_tensor@core@ffi:array_get( W_arr, W_idx ) end ) end ), Sum + Bias_v end ) end ) end). -file("src/viva_tensor/nn/conv.gleam", 191). ?DOC(false). -spec pad_1d(list(float()), integer(), integer(), integer(), integer()) -> list(float()). pad_1d(Data, Batch, Channels, Length, Padding) -> case Padding =:= 0 of true -> Data; false -> Arr = viva_tensor@core@ffi:list_to_array(Data), Padded_length = Length + (2 * Padding), _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, Padded_length - 1), gleam@list:map( _pipe@2, fun(I) -> Src = I - Padding, case (Src >= 0) andalso (Src < Length) of true -> viva_tensor@core@ffi:array_get( Arr, Base + Src ); false -> +0.0 end end ) end ) end ) end. -file("src/viva_tensor/nn/conv.gleam", 126). ?DOC(false). -spec conv1d_forward(conv1d_config(), viva_tensor@tensor:tensor()) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. conv1d_forward(Config, Input) -> In_shape = viva_tensor@tensor:shape(Input), case In_shape of [Batch, In_c, Length] when In_c =:= erlang:element(2, Config) -> Kernel_size = erlang:element(4, Config), Stride = erlang:element(5, Config), Padding = erlang:element(6, Config), 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) andalso (Stride > 0) of false -> {error, {invalid_shape, <<<<<<<<<<<<<<<<<<"conv1d_forward: 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>>}}; true -> In_data = viva_tensor@tensor:to_list(Input), W_data = viva_tensor@tensor:to_list( erlang:element(7, Config) ), Bias_data = case erlang:element(8, Config) of {some, B} -> viva_tensor@tensor:to_list(B); none -> gleam@list:repeat(+0.0, erlang:element(3, Config)) end, Padded = pad_1d(In_data, Batch, In_c, Length, Padding), Out = conv1d_compute( Padded, W_data, Bias_data, Batch, In_c, Padded_length, erlang:element(3, Config), Kernel_size, Stride, Out_length ), viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, erlang:element(3, Config), Out_length] ) end; _ -> {error, {invalid_shape, <<<<<<"conv1d_forward: input must have shape [batch, "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>/binary, ", length], got "/utf8>>/binary, (shape_to_string(In_shape))/binary>>}} end. -file("src/viva_tensor/nn/conv.gleam", 278). ?DOC(false). -spec conv3d_init( integer(), integer(), {integer(), integer(), integer()}, {integer(), integer(), integer()}, {integer(), integer(), integer()} ) -> conv3d_config(). conv3d_init(In_channels, Out_channels, Kernel_size, Stride, Padding) -> {Kd, Kh, Kw} = Kernel_size, {conv3d_config, In_channels, Out_channels, Kernel_size, Stride, Padding, viva_tensor@tensor:zeros([Out_channels, In_channels, Kd, Kh, Kw]), {some, viva_tensor@tensor:zeros([Out_channels])}}. -file("src/viva_tensor/nn/conv.gleam", 444). ?DOC(false). -spec pad_3d_internal( list(float()), integer(), integer(), integer(), integer(), integer(), integer(), integer(), integer() ) -> list(float()). pad_3d_internal(Data, Batch, Channels, Depth, Height, Width, Pd, Ph, Pw) -> case ((Pd =:= 0) andalso (Ph =:= 0)) andalso (Pw =:= 0) of true -> Data; false -> Arr = viva_tensor@core@ffi:list_to_array(Data), Padded_d = Depth + (2 * Pd), Padded_h = Height + (2 * Ph), Padded_w = Width + (2 * Pw), _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) * Depth) * Height) * Width) + (((C * Depth) * Height) * Width), _pipe@2 = range_int(0, Padded_d - 1), gleam@list:flat_map( _pipe@2, fun(Z) -> _pipe@3 = range_int(0, Padded_h - 1), gleam@list:flat_map( _pipe@3, fun(Y) -> _pipe@4 = range_int(0, Padded_w - 1), gleam@list:map( _pipe@4, fun(X) -> Sz = Z - Pd, Sy = Y - Ph, Sx = X - Pw, case (((((Sz >= 0) andalso (Sz < Depth)) andalso (Sy >= 0)) andalso (Sy < Height)) andalso (Sx >= 0)) andalso (Sx < Width) of true -> viva_tensor@core@ffi:array_get( Arr, ((Base + ((Sz * Height) * Width)) + (Sy * Width)) + Sx ); false -> +0.0 end end ) end ) end ) end ) end ) end. -file("src/viva_tensor/nn/conv.gleam", 310). ?DOC(false). -spec conv3d_forward(conv3d_config(), viva_tensor@tensor:tensor()) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. conv3d_forward(Config, Input) -> In_shape = viva_tensor@tensor:shape(Input), case In_shape of [Batch, In_c, Depth, Height, Width] when In_c =:= erlang:element( 2, Config ) -> {Kd, Kh, Kw} = erlang:element(4, Config), {Sd, Sh, Sw} = erlang:element(5, Config), {Pd, Ph, Pw} = erlang:element(6, Config), Padded_d = Depth + (2 * Pd), Padded_h = Height + (2 * Ph), Padded_w = Width + (2 * Pw), Out_d = (case Sd of 0 -> 0; Gleam@denominator -> (Padded_d - Kd) div Gleam@denominator end) + 1, Out_h = (case Sh of 0 -> 0; Gleam@denominator@1 -> (Padded_h - Kh) div Gleam@denominator@1 end) + 1, Out_w = (case Sw of 0 -> 0; Gleam@denominator@2 -> (Padded_w - Kw) div Gleam@denominator@2 end) + 1, case (((((Out_d > 0) andalso (Out_h > 0)) andalso (Out_w > 0)) andalso (Sd > 0)) andalso (Sh > 0)) andalso (Sw > 0) of false -> {error, {invalid_shape, <<<<<<<<<<<<"conv3d_forward: invalid output dims ("/utf8, (erlang:integer_to_binary( Out_d ))/binary>>/binary, ", "/utf8>>/binary, (erlang:integer_to_binary(Out_h))/binary>>/binary, ", "/utf8>>/binary, (erlang:integer_to_binary(Out_w))/binary>>/binary, ")"/utf8>>}}; true -> In_data = viva_tensor@tensor:to_list(Input), W_data = viva_tensor@tensor:to_list( erlang:element(7, Config) ), Bias_data = case erlang:element(8, Config) of {some, B} -> viva_tensor@tensor:to_list(B); none -> gleam@list:repeat(+0.0, erlang:element(3, Config)) end, Padded = pad_3d_internal( In_data, Batch, In_c, Depth, Height, Width, Pd, Ph, Pw ), In_arr = viva_tensor@core@ffi:list_to_array(Padded), W_arr = viva_tensor@core@ffi:list_to_array(W_data), Bias_arr = viva_tensor@core@ffi:list_to_array(Bias_data), Out = begin _pipe = range_int(0, Batch - 1), gleam@list:flat_map( _pipe, fun(B@1) -> _pipe@1 = range_int( 0, erlang:element(3, Config) - 1 ), gleam@list:flat_map( _pipe@1, fun(Oc) -> Bias_v = viva_tensor@core@ffi:array_get( Bias_arr, Oc ), _pipe@2 = range_int(0, Out_d - 1), gleam@list:flat_map( _pipe@2, fun(Od) -> _pipe@3 = range_int( 0, Out_h - 1 ), gleam@list:flat_map( _pipe@3, fun(Oh) -> _pipe@4 = range_int( 0, Out_w - 1 ), gleam@list:map( _pipe@4, fun(Ow) -> Sd_start = Od * Sd, Sh_start = Oh * Sh, Sw_start = Ow * Sw, Sum = sum_over_range( 0, In_c - 1, fun(Ic) -> sum_over_range( 0, Kd - 1, fun( Zk ) -> sum_over_range( 0, Kh - 1, fun( Yk ) -> sum_over_range( 0, Kw - 1, fun( Xk ) -> In_idx = ((((((((B@1 * In_c) * Padded_d) * Padded_h) * Padded_w) + (((Ic * Padded_d) * Padded_h) * Padded_w)) + (((Sd_start + Zk) * Padded_h) * Padded_w)) + ((Sh_start + Yk) * Padded_w)) + Sw_start) + Xk, W_idx = (((((((Oc * In_c) * Kd) * Kh) * Kw) + (((Ic * Kd) * Kh) * Kw)) + ((Zk * Kh) * Kw)) + (Yk * Kw)) + Xk, viva_tensor@core@ffi:array_get( In_arr, In_idx ) * viva_tensor@core@ffi:array_get( W_arr, W_idx ) end ) end ) end ) end ), Sum + Bias_v end ) end ) end ) end ) end ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, erlang:element(3, Config), Out_d, Out_h, Out_w] ) end; _ -> {error, {invalid_shape, <<<<<<"conv3d_forward: input must have shape [batch, "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>/binary, ", depth, height, width], got "/utf8>>/binary, (shape_to_string(In_shape))/binary>>}} end. -file("src/viva_tensor/nn/conv.gleam", 520). ?DOC(false). -spec conv_transpose_2d_init( integer(), integer(), {integer(), integer()}, {integer(), integer()}, {integer(), integer()}, {integer(), integer()} ) -> conv_transpose2d_config(). conv_transpose_2d_init( In_channels, Out_channels, Kernel_size, Stride, Padding, Output_padding ) -> {Kh, Kw} = Kernel_size, {conv_transpose2d_config, In_channels, Out_channels, Kernel_size, Stride, Padding, Output_padding, viva_tensor@tensor:zeros([In_channels, Out_channels, Kh, Kw]), {some, viva_tensor@tensor:zeros([Out_channels])}}. -file("src/viva_tensor/nn/conv.gleam", 662). ?DOC(false). -spec scatter_transpose_2d( viva_tensor@core@ffi:erlang_array(), viva_tensor@core@ffi:erlang_array(), integer(), integer(), integer(), integer(), integer(), integer(), integer(), integer(), integer(), integer(), integer() ) -> list(float()). scatter_transpose_2d( In_arr, W_arr, Batch, In_c, Out_c, H_in, W_in, Kh, Kw, Sh, Sw, H_full, W_full ) -> _pipe = range_int(0, Batch - 1), gleam@list:flat_map(_pipe, fun(B) -> _pipe@1 = range_int(0, Out_c - 1), gleam@list:flat_map( _pipe@1, fun(Oc) -> _pipe@2 = range_int(0, H_full - 1), gleam@list:flat_map( _pipe@2, fun(Y) -> _pipe@3 = range_int(0, W_full - 1), gleam@list:map( _pipe@3, fun(X) -> sum_over_range( 0, In_c - 1, fun(Ic) -> sum_over_range( 0, Kh - 1, fun(Ky) -> Dy = Y - Ky, case (Dy >= 0) andalso ((case Sh of 0 -> 0; Gleam@denominator -> Dy rem Gleam@denominator end) =:= 0) of false -> +0.0; true -> Ih = case Sh of 0 -> 0; Gleam@denominator@1 -> Dy div Gleam@denominator@1 end, case (Ih >= 0) andalso (Ih < H_in) of false -> +0.0; true -> sum_over_range( 0, Kw - 1, fun(Kx) -> Dx = X - Kx, case (Dx >= 0) andalso ((case Sw of 0 -> 0; Gleam@denominator@2 -> Dx rem Gleam@denominator@2 end) =:= 0) of false -> +0.0; true -> Iw = case Sw of 0 -> 0; Gleam@denominator@3 -> Dx div Gleam@denominator@3 end, case (Iw >= 0) andalso (Iw < W_in) of false -> +0.0; true -> In_idx = (((((B * In_c) * H_in) * W_in) + ((Ic * H_in) * W_in)) + (Ih * W_in)) + Iw, W_idx = (((((Ic * Out_c) * Kh) * Kw) + ((Oc * Kh) * Kw)) + (Ky * Kw)) + Kx, viva_tensor@core@ffi:array_get( In_arr, In_idx ) * viva_tensor@core@ffi:array_get( W_arr, W_idx ) end end end ) end end end ) end ) end ) end ) end ) end). -file("src/viva_tensor/nn/conv.gleam", 557). ?DOC(false). -spec conv_transpose_2d_forward( conv_transpose2d_config(), viva_tensor@tensor:tensor() ) -> {ok, viva_tensor@tensor:tensor()} | {error, viva_tensor@core@error:tensor_error()}. conv_transpose_2d_forward(Config, Input) -> In_shape = viva_tensor@tensor:shape(Input), case In_shape of [Batch, In_c, H_in, W_in] when In_c =:= erlang:element(2, Config) -> {Kh, Kw} = erlang:element(4, Config), {Sh, Sw} = erlang:element(5, Config), {Ph, Pw} = erlang:element(6, Config), {Oph, Opw} = erlang:element(7, Config), H_out = (((((H_in - 1) * Sh) - (2 * Ph)) + (Kh - 1)) + Oph) + 1, W_out = (((((W_in - 1) * Sw) - (2 * Pw)) + (Kw - 1)) + Opw) + 1, case (((H_out > 0) andalso (W_out > 0)) andalso (Sh > 0)) andalso (Sw > 0) of false -> {error, {invalid_shape, <<<<<<<<"conv_transpose_2d_forward: invalid output dims ("/utf8, (erlang:integer_to_binary(H_out))/binary>>/binary, ", "/utf8>>/binary, (erlang:integer_to_binary(W_out))/binary>>/binary, ")"/utf8>>}}; true -> In_data = viva_tensor@tensor:to_list(Input), W_data = viva_tensor@tensor:to_list( erlang:element(8, Config) ), Bias_data = case erlang:element(9, Config) of {some, B} -> viva_tensor@tensor:to_list(B); none -> gleam@list:repeat(+0.0, erlang:element(3, Config)) end, In_arr = viva_tensor@core@ffi:list_to_array(In_data), W_arr = viva_tensor@core@ffi:list_to_array(W_data), Bias_arr = viva_tensor@core@ffi:list_to_array(Bias_data), Out_c = erlang:element(3, Config), H_full = ((H_in - 1) * Sh) + Kh, W_full = ((W_in - 1) * Sw) + Kw, Acc = scatter_transpose_2d( In_arr, W_arr, Batch, In_c, Out_c, H_in, W_in, Kh, Kw, Sh, Sw, H_full, W_full ), Acc_arr = viva_tensor@core@ffi:list_to_array(Acc), Out = begin _pipe = range_int(0, Batch - 1), gleam@list:flat_map( _pipe, fun(B@1) -> _pipe@1 = range_int(0, Out_c - 1), gleam@list:flat_map( _pipe@1, fun(Oc) -> Bias_v = viva_tensor@core@ffi:array_get( Bias_arr, Oc ), _pipe@2 = range_int(0, H_out - 1), gleam@list:flat_map( _pipe@2, fun(Y) -> _pipe@3 = range_int( 0, W_out - 1 ), gleam@list:map( _pipe@3, fun(X) -> Src_y = Y + Ph, Src_x = X + Pw, case (((Src_y >= 0) andalso (Src_y < H_full)) andalso (Src_x >= 0)) andalso (Src_x < W_full) of true -> Idx = (((((B@1 * Out_c) * H_full) * W_full) + ((Oc * H_full) * W_full)) + (Src_y * W_full)) + Src_x, viva_tensor@core@ffi:array_get( Acc_arr, Idx ) + Bias_v; false -> Bias_v end end ) end ) end ) end ) end, viva_tensor@tensor:reshape( viva_tensor@tensor:from_list(Out), [Batch, Out_c, H_out, W_out] ) end; _ -> {error, {invalid_shape, <<<<<<"conv_transpose_2d_forward: input must have shape [batch, "/utf8, (erlang:integer_to_binary( erlang:element(2, Config) ))/binary>>/binary, ", height, width], got "/utf8>>/binary, (shape_to_string(In_shape))/binary>>}} end.