use crate::gpu::gpu_shader::moe_shader;
use wgpu::util::DeviceExt;
pub(super) struct MoeArtifacts {
pub(super) pipeline: wgpu::ComputePipeline,
pub(super) weights_buf: wgpu::Buffer,
pub(super) bind_group_layout: wgpu::BindGroupLayout,
}
pub(super) fn load_moe_artifacts(
device: &wgpu::Device,
adapter_info: &wgpu::AdapterInfo,
device_limits: &wgpu::Limits,
) -> Result<MoeArtifacts, String> {
let all_weights = crate::ml_scorer::ml_weights::all_weights_slice();
validate_weights_size(
std::mem::size_of_val(all_weights) as u64,
adapter_info,
device_limits,
)?;
let shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("moe_shader"),
source: wgpu::ShaderSource::Wgsl(moe_shader().into()),
});
let bind_group_layout = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor {
label: Some("moe_bgl"),
entries: &[
bgl_entry(0, true),
bgl_entry(1, true),
bgl_entry(2, false),
wgpu::BindGroupLayoutEntry {
binding: 3,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Uniform,
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
},
],
});
let pipeline_layout = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor {
label: Some("moe_pipeline_layout"),
bind_group_layouts: &[&bind_group_layout],
push_constant_ranges: &[],
});
let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("moe_pipeline"),
layout: Some(&pipeline_layout),
module: &shader,
entry_point: Some("moe_forward"),
compilation_options: Default::default(),
cache: None,
});
let weights_buf = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("weights"),
contents: bytemuck::cast_slice(all_weights),
usage: wgpu::BufferUsages::STORAGE,
});
Ok(MoeArtifacts {
pipeline,
weights_buf,
bind_group_layout,
})
}
pub(super) fn validate_weights_size(
weights_bytes: u64,
adapter_info: &wgpu::AdapterInfo,
device_limits: &wgpu::Limits,
) -> Result<(), String> {
let max_storage_binding = u64::from(device_limits.max_storage_buffer_binding_size);
if weights_bytes > max_storage_binding {
return Err(format!(
"GPU adapter {} exposes max_storage_buffer_binding_size={max_storage_binding} B, too small for the {weights_bytes} B MoE weights buffer",
adapter_info.name
));
}
Ok(())
}
fn bgl_entry(binding: u32, read_only: bool) -> wgpu::BindGroupLayoutEntry {
wgpu::BindGroupLayoutEntry {
binding,
visibility: wgpu::ShaderStages::COMPUTE,
ty: wgpu::BindingType::Buffer {
ty: wgpu::BufferBindingType::Storage { read_only },
has_dynamic_offset: false,
min_binding_size: None,
},
count: None,
}
}