#![allow(unused_imports)]
use crate::arena::Arena;
use crate::device::metal_device;
use crate::kernels::kernels;
use crate::thunk::{Thunk, ThunkSchedule};
use rlx_ir::{Graph, NodeId, Op};
use rlx_opt::memory;
use std::collections::HashMap;
use super::*;
impl MetalExecutable {
pub fn set_param(&mut self, name: &str, data: &[f32]) {
if let Some(&id) = self.param_ids.get(name) {
if let Some(slot) = self.weight_slots.get(&id).copied() {
self.write_weight_from_f32(slot, data);
} else if self.arena.has_buffer(id) {
self.arena.write_from_f32(id, data);
}
}
}
pub fn set_param_bytes(&mut self, name: &str, data: &[u8]) {
if let Some(&id) = self.param_ids.get(name) {
if let Some(slot) = self.weight_slots.get(&id).copied() {
self.write_weight_bytes(slot, data);
} else if self.arena.has_buffer(id) {
self.arena.write_bytes(id, data);
}
}
}
pub fn set_param_range(&mut self, name: &str, byte_offset: usize, data: &[u8]) -> bool {
let Some(&id) = self.param_ids.get(name) else {
return false;
};
if self.weight_slots.contains_key(&id) {
return false; }
if self.arena.has_buffer(id) {
self.arena.write_bytes_at(id, byte_offset, data);
return true;
}
false
}
pub fn param_storage_is_f16(&self, name: &str) -> bool {
let Some(&id) = self.param_ids.get(name) else {
return false;
};
if let Some(slot) = self.weight_slots.get(&id) {
return slot.dtype == rlx_ir::DType::F16;
}
self.arena.dtype(id) == rlx_ir::DType::F16
}
pub fn set_active_extent(&mut self, extent: Option<(usize, usize)>) {
self.active_extent = extent;
}
pub fn set_rng(&mut self, rng: rlx_ir::RngOptions) {
*self.schedule.rng.write().expect("rng lock") = rng;
}
pub fn set_gpu_handle_feed(&mut self, handle_name: &str, output_index: usize) {
self.gpu_handle_feeds
.insert(handle_name.to_string(), output_index);
}
}