Files
rustytorch/crates/models/rtx-diffuse/benches/diffusion_bench.rs
T
2026-03-04 00:08:42 +00:00

349 lines
9.6 KiB
Rust

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(&timesteps))
.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(&timesteps),
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(&timesteps)).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(&timesteps))
.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);