use scirs2_core::ndarray::{Array, IxDyn};
use scirs2_neural::error::Result;
use scirs2_neural::layers::Layer;
use scirs2_neural::models::{ViTConfig, VisionTransformer};
fn main() -> Result<()> {
println!("Vision Transformer (ViT) Example");
println!("================================");
let config = ViTConfig {
image_size: (32, 32), patch_size: (8, 8), in_channels: 3, num_classes: 10, embed_dim: 32, num_layers: 2, num_heads: 4, mlp_dim: 64, dropout_rate: 0.1,
attention_dropout_rate: 0.1,
};
println!(
"Creating custom ViT model: image {:?}, patch {:?}, {} channels, {} classes",
config.image_size, config.patch_size, config.in_channels, config.num_classes
);
let model = VisionTransformer::<f32>::new(config)?;
let input = Array::from_shape_fn(IxDyn(&[1, 3, 32, 32]), |_| {
scirs2_core::random::random::<f32>()
});
println!("Input shape: {:?}", input.shape());
let output = model.forward(&input)?;
println!("Output shape: {:?}", output.shape());
println!("Output contains logits for {} classes", output.shape()[1]);
println!("\nCreating a ViT-Base model...");
let base_model = VisionTransformer::<f32>::vit_base((224, 224), (16, 16), 3, 1000)?;
let base_config = base_model.config();
println!("ViT-Base model created successfully.");
println!(" - embedding dimension: {}", base_config.embed_dim);
println!(" - transformer layers: {}", base_config.num_layers);
println!(" - attention heads: {}", base_config.num_heads);
println!("\nVision Transformer example completed successfully!");
Ok(())
}