use serde::{Deserialize, Serialize};
use vyre_lower::KernelDescriptor;
pub const VULKAN_BASELINE: DeviceLimits = DeviceLimits {
max_workgroup_size_per_dim: [1024, 1024, 64],
max_workgroup_invocations: 1024,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct DeviceLimits {
pub max_workgroup_size_per_dim: [u32; 3],
pub max_workgroup_invocations: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum Violation {
DimExceeded { axis: u8, actual: u32, limit: u32 },
InvocationsExceeded { actual: u32, limit: u32 },
ZeroDim { axis: u8 },
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct ValidationReport {
pub kernel_id: String,
pub workgroup_size: [u32; 3],
pub limits: DeviceLimits,
pub violations: Vec<Violation>,
}
impl ValidationReport {
pub fn ok(&self) -> bool {
self.violations.is_empty()
}
pub fn invocations(&self) -> u32 {
self.workgroup_size[0]
.saturating_mul(self.workgroup_size[1])
.saturating_mul(self.workgroup_size[2])
}
}
#[must_use]
pub fn analyze(desc: &KernelDescriptor) -> ValidationReport {
analyze_against(desc, VULKAN_BASELINE)
}
#[must_use]
pub fn analyze_against(desc: &KernelDescriptor, limits: DeviceLimits) -> ValidationReport {
let wg = desc.dispatch.workgroup_size;
let mut violations = Vec::new();
for (axis, &dim) in wg.iter().enumerate() {
let axis = axis as u8;
if dim == 0 {
violations.push(Violation::ZeroDim { axis });
} else if dim > limits.max_workgroup_size_per_dim[axis as usize] {
violations.push(Violation::DimExceeded {
axis,
actual: dim,
limit: limits.max_workgroup_size_per_dim[axis as usize],
});
}
}
let invocations = wg[0].saturating_mul(wg[1]).saturating_mul(wg[2]);
if invocations > limits.max_workgroup_invocations {
violations.push(Violation::InvocationsExceeded {
actual: invocations,
limit: limits.max_workgroup_invocations,
});
}
ValidationReport {
kernel_id: desc.id.clone(),
workgroup_size: wg,
limits,
violations,
}
}
#[cfg(test)]
mod tests {
use super::*;
use vyre_lower::{BindingLayout, Dispatch, KernelBody, KernelDescriptor};
fn empty_with_dispatch(d: Dispatch) -> KernelDescriptor {
KernelDescriptor {
id: "k".into(),
bindings: BindingLayout { slots: vec![] },
dispatch: d,
body: KernelBody {
ops: vec![],
child_bodies: vec![],
literals: vec![],
},
}
}
#[test]
fn small_workgroup_is_valid() {
let report = analyze(&empty_with_dispatch(Dispatch::new(64, 1, 1)));
assert!(report.ok());
assert_eq!(report.invocations(), 64);
}
#[test]
fn standard_1d_1024_workgroup_is_valid_at_baseline() {
let report = analyze(&empty_with_dispatch(Dispatch::new(1024, 1, 1)));
assert!(report.ok());
}
#[test]
fn dim_x_over_1024_violates_dim_limit() {
let report = analyze(&empty_with_dispatch(Dispatch::new(2048, 1, 1)));
assert!(!report.ok());
let has_dim_violation = report
.violations
.iter()
.any(|v| matches!(v, Violation::DimExceeded { axis: 0, .. }));
assert!(has_dim_violation);
}
#[test]
fn dim_z_over_64_violates_baseline() {
let report = analyze(&empty_with_dispatch(Dispatch::new(1, 1, 128)));
assert!(!report.ok());
let has = report.violations.iter().any(|v| {
matches!(
v,
Violation::DimExceeded {
axis: 2,
actual: 128,
limit: 64
}
)
});
assert!(has);
}
#[test]
fn product_over_1024_violates_invocations() {
let report = analyze(&empty_with_dispatch(Dispatch::new(32, 32, 2)));
assert!(!report.ok());
let has = report
.violations
.iter()
.any(|v| matches!(v, Violation::InvocationsExceeded { actual: 2048, .. }));
assert!(has);
}
#[test]
fn zero_dim_y_flagged() {
let report = analyze(&empty_with_dispatch(Dispatch::new(64, 0, 1)));
let has = report
.violations
.iter()
.any(|v| matches!(v, Violation::ZeroDim { axis: 1 }));
assert!(has);
}
#[test]
fn high_end_device_profile_allows_more() {
let limits = DeviceLimits {
max_workgroup_size_per_dim: [1024, 1024, 1024],
max_workgroup_invocations: 1024,
};
let report = analyze_against(&empty_with_dispatch(Dispatch::new(1, 1, 128)), limits);
assert!(report.ok());
}
#[test]
fn invocations_helper_computes_product() {
let report = analyze(&empty_with_dispatch(Dispatch::new(8, 8, 4)));
assert_eq!(report.invocations(), 256);
}
#[test]
fn carries_kernel_id() {
let mut desc = empty_with_dispatch(Dispatch::new(1, 1, 1));
desc.id = "named".into();
let report = analyze(&desc);
assert_eq!(report.kernel_id, "named");
}
}