use criterion::{BenchmarkId, Criterion, black_box, criterion_group, criterion_main}; use rtx_diffuse::{ DiT, DiTConfig, DiffusionScheduler, NoiseGenerator, NoiseSchedule, Result, SchedulerType, UNet, UNetConfig, }; use rtx_tensor::Tensor; fn bench_noise_generation(c: &mut Criterion) { let mut group = c.benchmark_group("noise_generation"); let mut noise_gen = NoiseGenerator::new( NoiseSchedule::Linear { beta_start: 0.0001, beta_end: 0.02, }, 1000, Some(42), ) .unwrap(); for size in [16, 32, 64].iter() { group.bench_with_input( BenchmarkId::new("generate_noise", size), size, |b, &size| { b.iter(|| { let noise = noise_gen .generate_noise(black_box(&[2, 4, size, size])) .unwrap(); black_box(noise); }); }, ); } group.finish(); } fn bench_noise_schedules(c: &mut Criterion) { let mut group = c.benchmark_group("noise_schedules"); let schedules = vec![ ( "linear", NoiseSchedule::Linear { beta_start: 0.0001, beta_end: 0.02, }, ), ("cosine", NoiseSchedule::Cosine { s: 0.008 }), ( "scaled_linear", NoiseSchedule::ScaledLinear { beta_start: 0.00085, beta_end: 0.012, }, ), ]; for (name, schedule) in schedules { group.bench_function(name, |b| { b.iter(|| { let noise_gen = NoiseGenerator::new(black_box(schedule.clone()), black_box(1000), Some(42)) .unwrap(); black_box(noise_gen); }); }); } group.finish(); } fn bench_add_noise(c: &mut Criterion) { let mut group = c.benchmark_group("add_noise"); let noise_gen = NoiseGenerator::new( NoiseSchedule::Linear { beta_start: 0.0001, beta_end: 0.02, }, 1000, Some(42), ) .unwrap(); for size in [16, 32].iter() { let x0 = Tensor::new(vec![0.5; 2 * 4 * size * size], vec![2, 4, *size, *size]).unwrap(); let noise = Tensor::new(vec![0.1; 2 * 4 * size * size], vec![2, 4, *size, *size]).unwrap(); group.bench_with_input(BenchmarkId::new("add_noise", size), size, |b, _| { b.iter(|| { let noisy = noise_gen .add_noise(black_box(&x0), black_box(&noise), black_box(500)) .unwrap(); black_box(noisy); }); }); } group.finish(); } fn bench_scheduler_step(c: &mut Criterion) { let mut group = c.benchmark_group("scheduler_step"); let noise_gen = NoiseGenerator::new( NoiseSchedule::Linear { beta_start: 0.0001, beta_end: 0.02, }, 1000, Some(42), ) .unwrap(); let schedulers = vec![ ("ddim", SchedulerType::DDIM { eta: 0.0 }), ("dpm_plusplus", SchedulerType::DPMPlusPlus), ("euler_ancestral", SchedulerType::EulerAncestral), ("ddpm", SchedulerType::DDPM), ]; for (name, scheduler_type) in schedulers { let scheduler = DiffusionScheduler::new(scheduler_type, noise_gen.clone(), 50).unwrap(); let sample = Tensor::new(vec![0.3; 1 * 4 * 32 * 32], vec![1, 4, 32, 32]).unwrap(); let model_output = Tensor::new(vec![0.1; 1 * 4 * 32 * 32], vec![1, 4, 32, 32]).unwrap(); let timestep = scheduler.timesteps()[10]; group.bench_function(name, |b| { b.iter(|| { let result = scheduler .step( black_box(&model_output), black_box(timestep), black_box(&sample), None, ) .unwrap(); black_box(result); }); }); } group.finish(); } fn bench_unet_forward(c: &mut Criterion) { let mut group = c.benchmark_group("unet_forward"); let configs = vec![ ( "small", UNetConfig { model_channels: 128, num_res_blocks: 1, channel_mult: vec![1, 2], ..Default::default() }, ), ("default", UNetConfig::default()), ]; for (name, config) in configs { let unet = UNet::new(config).unwrap(); let input = Tensor::new(vec![0.5; 1 * 4 * 32 * 32], vec![1, 4, 32, 32]).unwrap(); let timesteps = Tensor::new(vec![500.0], vec![1]).unwrap(); group.bench_function(name, |b| { b.iter(|| { let output = unet .forward(black_box(&input), black_box(×teps)) .unwrap(); black_box(output); }); }); } group.finish(); } fn bench_dit_forward(c: &mut Criterion) { let mut group = c.benchmark_group("dit_forward"); let configs = vec![ ( "small", DiTConfig { input_size: 16, patch_size: 2, hidden_size: 192, depth: 4, num_heads: 8, ..Default::default() }, ), ( "tiny", DiTConfig { input_size: 8, patch_size: 2, hidden_size: 96, depth: 2, num_heads: 4, ..Default::default() }, ), ]; for (name, config) in configs { let dit = DiT::new(config.clone()).unwrap(); let input_size = config.input_size; let input = Tensor::new( vec![0.3; 1 * 4 * input_size * input_size], vec![1, 4, input_size, input_size], ) .unwrap(); let timesteps = Tensor::new(vec![400.0], vec![1]).unwrap(); let labels = Tensor::new(vec![5.0], vec![1]).unwrap(); group.bench_function(name, |b| { b.iter(|| { let output = dit .forward( black_box(&input), black_box(×teps), black_box(Some(&labels)), ) .unwrap(); black_box(output); }); }); } group.finish(); } fn bench_time_embedding(c: &mut Criterion) { let mut group = c.benchmark_group("time_embedding"); for batch_size in [1, 4, 8].iter() { group.bench_with_input( BenchmarkId::new("unet_time_embedding", batch_size), batch_size, |b, &batch_size| { use rtx_diffuse::models::unet::TimeEmbedding; let time_emb = TimeEmbedding::new(512).unwrap(); let timesteps = Tensor::new(vec![500.0; batch_size], vec![batch_size]).unwrap(); b.iter(|| { let embeddings = time_emb.forward(black_box(×teps)).unwrap(); black_box(embeddings); }); }, ); } group.finish(); } fn bench_patch_embedding(c: &mut Criterion) { let mut group = c.benchmark_group("patch_embedding"); use rtx_diffuse::models::dit::PatchEmbed; for img_size in [16, 32].iter() { let patch_embed = PatchEmbed::new(*img_size, 2, 4, 384).unwrap(); let input = Tensor::new( vec![0.2; 2 * 4 * img_size * img_size], vec![2, 4, *img_size, *img_size], ) .unwrap(); group.bench_with_input( BenchmarkId::new("patch_embed", img_size), img_size, |b, _| { b.iter(|| { let patches = patch_embed.forward(black_box(&input)).unwrap(); black_box(patches); }); }, ); } group.finish(); } fn bench_full_diffusion_step(c: &mut Criterion) { let mut group = c.benchmark_group("full_diffusion_step"); let noise_gen = NoiseGenerator::new(NoiseSchedule::Cosine { s: 0.008 }, 1000, Some(42)).unwrap(); let scheduler = DiffusionScheduler::new(SchedulerType::DDIM { eta: 0.0 }, noise_gen, 50).unwrap(); let unet_config = UNetConfig { model_channels: 64, num_res_blocks: 1, channel_mult: vec![1, 2], ..Default::default() }; let unet = UNet::new(unet_config).unwrap(); let sample = Tensor::new(vec![0.5; 1 * 4 * 16 * 16], vec![1, 4, 16, 16]).unwrap(); let timesteps = Tensor::new(vec![scheduler.timesteps()[10] as f32], vec![1]).unwrap(); group.bench_function("unet_plus_scheduler", |b| { b.iter(|| { // UNet forward pass let model_output = unet .forward(black_box(&sample), black_box(×teps)) .unwrap(); // Scheduler step let result = scheduler .step( black_box(&model_output), black_box(scheduler.timesteps()[10]), black_box(&sample), None, ) .unwrap(); black_box(result); }); }); group.finish(); } criterion_group!( benches, bench_noise_generation, bench_noise_schedules, bench_add_noise, bench_scheduler_step, bench_unet_forward, bench_dit_forward, bench_time_embedding, bench_patch_embedding, bench_full_diffusion_step, ); criterion_main!(benches);