1pub mod activation;
24pub mod conv;
25pub mod core;
26pub mod linear;
27pub mod loss;
28pub mod loss_advanced;
29pub mod norm;
30pub mod pooling;
31
32pub use core::*;
34
35pub use activation::{
41 dropout, elu, gelu, leaky_relu, log_softmax, mish, relu, relu_inplace, selu, sigmoid, softmax,
42 swish, tanh,
43};
44
45pub use activation::{
47 LeakyReLU, LogSoftmax, Mish, ReLU, Sigmoid, Softmax, Swish, Tanh, ELU, GELU, SELU,
48};
49
50pub use activation::configured::{
52 gelu_configured, mish_configured, relu_configured, sigmoid_configured, softmax_configured,
53 swish_configured, tanh_configured,
54};
55
56pub use conv::{conv1d, conv2d, conv3d, conv_transpose1d, conv_transpose2d, conv_transpose3d};
62
63pub use conv::{conv_output_size, conv_transpose_output_size, validate_conv_params};
65
66pub use pooling::{
72 adaptive_max_pool1d, adaptive_max_pool2d, adaptive_max_pool3d, global_max_pool1d,
73 global_max_pool2d, global_max_pool3d, max_pool1d, max_pool2d, max_pool3d,
74};
75
76pub use pooling::{
78 adaptive_avg_pool1d, adaptive_avg_pool2d, adaptive_avg_pool3d, avg_pool1d, avg_pool2d,
79 avg_pool3d, global_avg_pool1d, global_avg_pool2d, global_avg_pool3d,
80};
81
82pub use pooling::{
84 pad, reflection_pad1d, reflection_pad2d, replication_pad1d, replication_pad2d, zero_pad2d,
85};
86
87pub use pooling::{adaptive_pool_params, pool_output_size};
89
90pub use linear::{bilinear, linear};
96
97pub use linear::{embedding, embedding_bag, one_hot};
99
100pub use linear::{
102 grouped_query_attention, multi_head_attention, multi_query_attention,
103 scaled_dot_product_attention,
104};
105
106pub use linear::{
108 learnable_positional_encoding, rotary_positional_encoding, sinusoidal_positional_encoding,
109};
110
111pub use linear::{post_norm_layer_norm, pre_norm_layer_norm, rms_norm};
113
114pub use linear::{geglu, glu, swiglu};
116
117pub use linear::{apply_attention_mask, create_causal_mask, create_padding_mask};
119
120pub use loss::{
126 binary_cross_entropy, binary_cross_entropy_with_logits, cross_entropy, focal_loss,
127 multi_margin_loss, multilabel_margin_loss, nll_loss,
128};
129
130pub use loss::{huber_loss, l1_loss, mse_loss, smooth_l1_loss};
132
133pub use loss::kl_div;
135
136pub use loss::{contrastive_loss, cosine_embedding_loss, triplet_margin_loss};
138
139pub use loss::{center_loss, dice_loss, infonce_loss, tversky_loss, wing_loss};
141
142pub use loss_advanced::{CustomLoss, LossBuilder, LossFactory, Reduction};
148
149pub use loss_advanced::{
151 AdaptiveLoss, CombinedLoss, DiceLoss, IoULoss, SmoothL1Loss, WeightedLoss,
152};
153
154pub use loss_advanced::{
156 BinaryCrossEntropy, CategoricalCrossEntropy, CosineEmbeddingLoss, FocalLoss, HingeLoss,
157 HuberLoss, KLDivLoss, L1Loss, MSELoss, NLLLoss, TripletMarginLoss,
158};
159
160pub use loss_advanced::validation as loss_validation;
162
163pub use norm::{
169 batch_norm, batch_norm_1d, batch_norm_2d, batch_norm_2d_with_config, batch_norm_3d,
170};
171
172pub use norm::{layer_norm, layer_norm_configured, layer_norm_enhanced};
174
175pub use norm::{
177 group_norm, instance_norm, local_response_norm, rms_norm as rms_norm_standalone, spectral_norm,
178 weight_norm,
179};
180
181pub use norm::configured::batch_norm_configured;
183
184pub use norm::{create_affine_params, get_norm_features, validate_norm_params};
186
187pub mod activations {
193 pub use super::activation::configured::*;
194 pub use super::activation::{gelu, mish, relu, sigmoid, softmax, swish, tanh};
196}
197
198pub mod losses {
200 pub use super::loss::*;
201
202 pub fn mse_loss_configured(
204 input: &crate::Tensor,
205 target: &crate::Tensor,
206 reduction: &str,
207 config: &super::FunctionalConfig,
208 ) -> super::FuncResult<crate::Tensor> {
209 crate::validate_inputs!(
210 config,
211 super::validation::validate_not_empty(input, "input"),
212 super::validation::validate_not_empty(target, "target"),
213 super::validation::validate_compatible_shapes(input, target, "MSE loss")
214 );
215 crate::func_error!(super::mse_loss(input, target, reduction), "MSE loss")
216 }
217
218 pub fn l1_loss_configured(
220 input: &crate::Tensor,
221 target: &crate::Tensor,
222 reduction: &str,
223 config: &super::FunctionalConfig,
224 ) -> super::FuncResult<crate::Tensor> {
225 crate::validate_inputs!(
226 config,
227 super::validation::validate_not_empty(input, "input"),
228 super::validation::validate_not_empty(target, "target"),
229 super::validation::validate_compatible_shapes(input, target, "L1 loss")
230 );
231 crate::func_error!(super::l1_loss(input, target, reduction), "L1 loss")
232 }
233
234 pub fn cross_entropy_configured(
236 input: &crate::Tensor,
237 target: &crate::Tensor<i64>,
238 weight: Option<&crate::Tensor>,
239 ignore_index: Option<i64>,
240 reduction: &str,
241 config: &super::FunctionalConfig,
242 ) -> super::FuncResult<crate::Tensor> {
243 crate::validate_inputs!(
244 config,
245 super::validation::validate_not_empty(input, "input"),
246 super::validation::validate_not_empty(target, "target"),
247 super::validation::validate_min_ndim(input, 2, "input")
248 );
249 crate::func_error!(
250 super::cross_entropy(input, target, weight, reduction, ignore_index),
251 "Cross entropy loss"
252 )
253 }
254}
255
256pub mod normalization {
258 pub use super::norm::configured::*;
259 pub use super::norm::{batch_norm_2d, layer_norm_enhanced};
261}
262
263pub mod prelude {
269 pub use super::{
270 activations, default_config, losses, normalization, numerics, optimized, performance, safe,
271 validation, Activation, ActivationConfig, CustomLoss, FunctionalBuilder, FunctionalConfig,
272 LossBuilder, MemoryOptLevel, Reduction,
273 };
274}
275
276use torsh_core::error::Result;
282use torsh_tensor::Tensor;
283
284#[allow(dead_code)]
286trait TensorCast {
287 fn cast_i64(&self) -> Result<Tensor<i64>>;
288}
289
290#[allow(dead_code)]
291impl TensorCast for Tensor {
292 fn cast_i64(&self) -> Result<Tensor<i64>> {
293 let data = self.to_vec()?;
295 let i64_data: Vec<i64> = data.into_iter().map(|x| x as i64).collect();
296 Ok(Tensor::from_data(
297 i64_data,
298 self.shape().dims().to_vec(),
299 self.device(),
300 )?)
301 }
302}
303
304pub struct SparseMatrix;
306
307impl SparseMatrix {
308 pub fn new() -> Self {
309 Self
310 }
311}
312
313impl Default for SparseMatrix {
314 fn default() -> Self {
315 Self::new()
316 }
317}
318
319#[cfg(test)]
334mod tests {
335 use super::*;
336
337 #[test]
338 fn test_modular_functional_system() {
339 let input = torsh_tensor::creation::randn::<f32>(&[2, 4]).unwrap();
343 let _relu_result = relu(&input).unwrap();
344 let _sigmoid_result = sigmoid(&input).unwrap();
345 let _tanh_result = tanh(&input).unwrap();
346
347 let config = FunctionalConfig::default();
349 let _configured_relu = activations::relu_configured(&input, &config).unwrap();
350
351 let _optimized_config = optimized().build();
353 let _safe_config = safe().build();
354 }
355
356 #[test]
357 fn test_backward_compatibility() {
358 let input = torsh_tensor::creation::randn::<f32>(&[4, 3, 32, 32]).unwrap();
362 let weight = torsh_tensor::creation::ones(&[3]).unwrap();
363 let bias = torsh_tensor::creation::zeros(&[3]).unwrap();
364
365 let _batch_norm_result = batch_norm_2d(
367 &input,
368 Some(&weight),
369 Some(&bias),
370 None,
371 None,
372 true,
373 0.1,
374 1e-5,
375 )
376 .unwrap();
377
378 let activation_input = torsh_tensor::creation::randn::<f32>(&[2, 4]).unwrap();
380 let _relu_result = relu(&activation_input).unwrap();
381 let _gelu_result = gelu(&activation_input).unwrap();
382 let _swish_result = swish(&activation_input).unwrap();
383 }
384
385 #[test]
386 fn test_modular_structure_integrity() {
387 let config = FunctionalConfig::default();
391 assert_eq!(config.validate_inputs, true);
392 assert_eq!(config.eps, 1e-8);
393
394 let custom_config = FunctionalBuilder::new().eps(1e-6).inplace(true).build();
396 assert_eq!(custom_config.eps, 1e-6);
397 assert_eq!(custom_config.inplace, true);
398
399 let _default_conf = prelude::default_config();
401 }
402
403 #[test]
404 fn test_loss_framework() {
405 let predictions = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
407 let targets = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
408
409 let mse = MSELoss::new(Reduction::Mean);
411 let _loss_result = mse.compute_loss(&predictions, &targets).unwrap();
412
413 let _smooth_l1 = LossBuilder::new()
415 .with_reduction(Reduction::Sum)
416 .smooth_l1(1.0);
417 }
418}
419
420#[cfg(test)]
422mod examples {
423 use super::*;
424
425 #[test]
426 fn example_basic_usage() {
427 let input = torsh_tensor::creation::randn::<f32>(&[4, 3, 32, 32]).unwrap();
429 let target = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
430
431 let activated = relu(&input).unwrap();
433 let _softmax_result = softmax(&activated, Some(-1)).unwrap();
434
435 let weight = torsh_tensor::creation::ones(&[3]).unwrap();
437 let bias = torsh_tensor::creation::zeros(&[3]).unwrap();
438 let _normalized = batch_norm_2d(
439 &input,
440 Some(&weight),
441 Some(&bias),
442 None,
443 None,
444 true,
445 0.1,
446 1e-5,
447 )
448 .unwrap();
449
450 let predictions = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
452 let _mse_loss = mse_loss(&predictions, &target, "mean").unwrap();
453 }
454
455 #[test]
456 fn example_configured_usage() {
457 let config = FunctionalBuilder::new()
459 .validate(true)
460 .eps(1e-6)
461 .memory_opt(MemoryOptLevel::Maximum)
462 .build();
463
464 let input = torsh_tensor::creation::randn::<f32>(&[4, 8]).unwrap();
465
466 let _relu_result = activations::relu_configured(&input, &config).unwrap();
468 let _sigmoid_result = activations::sigmoid_configured(&input, &config).unwrap();
469 }
470
471 #[test]
472 fn example_advanced_loss_usage() {
473 let predictions = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
474 let targets = torsh_tensor::creation::randn::<f32>(&[4, 10]).unwrap();
475
476 let dice_loss = LossBuilder::new()
478 .with_reduction(Reduction::Mean)
479 .dice(1e-6);
480
481 let _loss_result = dice_loss.compute_loss(&predictions, &targets).unwrap();
482
483 let mse = Box::new(MSELoss::new(Reduction::None));
485 let l1 = Box::new(L1Loss::new(Reduction::None));
486
487 let combined = LossBuilder::new()
488 .with_reduction(Reduction::Mean)
489 .combined(vec![mse, l1], vec![0.7, 0.3]);
490
491 let _combined_loss = combined.compute_loss(&predictions, &targets).unwrap();
492 }
493}