use super::{Device, DeviceId, DeviceInfo, DeviceProgramId, Event, Kernel, MemoryPool, PoolBufferId, PoolId};
use crate::{
DType, Map,
backend::DTypeCapability,
error::{BackendError, ErrorStatus},
kernel::{Op, OpId, Scope},
shape::Dim,
slab::Slab,
};
use nanoserde::DeJson;
use std::{
ffi::CString,
io::{BufRead, BufReader, BufWriter, Write as IoWrite},
path::PathBuf,
process::{Child, ChildStdin, ChildStdout, Command},
sync::{Arc, Mutex},
};
const DRAM_SIZE_TABLE: &[(u16, &str, u64)] = &[
(0x0036, "p100", 28u64 * 1024 * 1024 * 1024),
(0x0040, "p150a", 32u64 * 1024 * 1024 * 1024),
(0x0041, "p150b", 32u64 * 1024 * 1024 * 1024),
(0x0042, "p150c", 32u64 * 1024 * 1024 * 1024),
(0x0043, "p100a", 28u64 * 1024 * 1024 * 1024),
(0x0044, "p300b", 64u64 * 1024 * 1024 * 1024),
(0x0045, "p300a", 64u64 * 1024 * 1024 * 1024),
(0x0046, "p300c", 64u64 * 1024 * 1024 * 1024),
];
fn detect_dram_bytes() -> u64 {
let pci_devices = std::path::Path::new("/sys/bus/pci/devices");
if let Ok(entries) = std::fs::read_dir(pci_devices) {
for entry in entries.flatten() {
let vendor_path = entry.path().join("vendor");
let vendor = std::fs::read_to_string(&vendor_path).unwrap_or_default();
if vendor.trim() == "0x1e52" {
let subsys = std::fs::read_to_string(entry.path().join("subsystem_device")).unwrap_or_default();
if let Ok(id) = u16::from_str_radix(subsys.trim().trim_start_matches("0x"), 16) {
for &(sid, _name, size) in DRAM_SIZE_TABLE {
if sid == id {
return size;
}
}
}
}
}
}
64u64 * 1024 * 1024 * 1024
}
#[derive(Default, Debug, DeJson)]
#[nserde(default)]
pub struct TTConfig {
device_ids: Option<Vec<i32>>,
}
#[derive(Debug)]
struct TTBuffer {
dev_index: u32,
size: u64,
}
#[derive(Debug)]
pub struct TTMemoryPool {
buffers: Slab<PoolBufferId, TTBuffer>,
runtime: Arc<Mutex<RuntimeProcess>>,
free_bytes: Dim,
}
#[derive(Debug, Clone)]
pub struct TTEvent;
pub(super) fn initialize_device(
config: &TTConfig,
memory_pools: &mut Slab<PoolId, MemoryPool>,
devices: &mut Slab<DeviceId, Device>,
debug_dev: bool,
) -> Result<(), BackendError> {
if let Some(device_ids) = &config.device_ids
&& device_ids.is_empty()
{
if debug_dev {
println!("[tenstorrent] configured out");
}
return Ok(());
}
let dram_bytes = detect_dram_bytes();
if debug_dev {
println!("[tenstorrent] device initialized");
println!("[tenstorrent] device total memory: {} MB", dram_bytes / (1024 * 1024));
}
let config_base = std::env::var_os("XDG_CONFIG_HOME")
.and_then(|p| {
let p = PathBuf::from(p);
if p.is_absolute() { Some(p) } else { None }
})
.or_else(|| std::env::home_dir().map(|h| h.join(".config")))
.unwrap_or_else(|| PathBuf::from("/tmp"));
let cache_dir = config_base.join("zyx/cache/tt");
let runtime_path = config_base.join("zyx/zyx-tt-runtime");
if !runtime_path.exists() {
return Err(BackendError {
status: ErrorStatus::Initialization,
context: format!("runtime not found at {}. Rebuild with TT_METAL_ROOT set.", runtime_path.display()).into(),
});
}
let runtime = Arc::new(Mutex::new(RuntimeProcess::new(&runtime_path.to_string_lossy(), &cache_dir.to_string_lossy())?));
let pool_id = memory_pools.len();
let pool = MemoryPool::TT(TTMemoryPool { buffers: Slab::new(), runtime: runtime.clone(), free_bytes: Dim::from(dram_bytes) });
memory_pools.push(pool);
let _device_id = devices.len();
devices.push(Device::TT(TTDevice {
device_info: DeviceInfo {
compute: 200_000_000_000_000, max_global_work_dims: vec![Dim::from(u32::MAX); 3],
max_local_threads: 1024,
max_local_work_dims: vec![1, 1024, 1],
preferred_vector_size: 32,
local_mem_size: 1_500_000, max_register_bytes: 128,
tensor_cores: true,
warp_size: 1, dtype_capability: [DTypeCapability::all(); DType::N_DTYPES],
has_native_exp2: false,
supported_vec_lens: vec![32],
},
memory_pool_id: pool_id,
runtime,
programs: Slab::new(),
}));
Ok(())
}
fn create_temp_shm(size: u64) -> Result<(CString, *mut u8, u64), BackendError> {
let pid = std::process::id();
let ns = std::time::SystemTime::now().duration_since(std::time::UNIX_EPOCH).unwrap_or_default().as_nanos();
let name = format!("/zyx-tt-{pid:x}-{ns:x}");
let cname = CString::new(name.clone())
.map_err(|_| BackendError { status: ErrorStatus::MemoryAllocation, context: "invalid shm path".into() })?;
let fd = unsafe { libc::shm_open(cname.as_ptr(), libc::O_CREAT | libc::O_RDWR | libc::O_EXCL, 0o600) };
if fd < 0 {
return Err(BackendError {
status: ErrorStatus::MemoryAllocation,
context: format!("shm_open errno={}", std::io::Error::last_os_error()).into(),
});
}
if unsafe { libc::ftruncate(fd, size as i64) } < 0 {
unsafe { libc::close(fd) };
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
return Err(BackendError { status: ErrorStatus::MemoryAllocation, context: "ftruncate shm".into() });
}
let ptr =
unsafe { libc::mmap(std::ptr::null_mut(), size as usize, libc::PROT_READ | libc::PROT_WRITE, libc::MAP_SHARED, fd, 0) };
if ptr == libc::MAP_FAILED {
unsafe { libc::close(fd) };
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
return Err(BackendError { status: ErrorStatus::MemoryAllocation, context: "mmap shm".into() });
}
unsafe { libc::close(fd) };
Ok((cname, ptr as *mut u8, size))
}
impl TTMemoryPool {
pub fn deinitialize(&mut self) {
let _ = self.runtime.lock().unwrap().exit();
}
pub fn free_bytes(&self) -> Dim {
self.free_bytes
}
pub fn allocate(&mut self, bytes: Dim) -> Result<(PoolBufferId, Event), BackendError> {
let bytes_u64: u64 = u64::try_from(bytes).map_err(|_| BackendError {
status: ErrorStatus::MemoryAllocation,
context: "allocation size exceeds 64-bit".into(),
})?;
if bytes > self.free_bytes {
return Err(BackendError { status: ErrorStatus::MemoryAllocation, context: "out of device memory".into() });
}
let rt = &self.runtime;
let tile_bytes: u64 = 2048;
let dev_index = rt.lock().unwrap().alloc_buf(bytes_u64, tile_bytes)?;
let buf = TTBuffer { dev_index, size: bytes_u64 };
let id = self.buffers.push(buf);
Ok((id, Event::TT(TTEvent)))
}
pub fn deallocate(&mut self, buffer_id: PoolBufferId, event_wait_list: Vec<Event>) {
let _ = event_wait_list;
if self.buffers.contains_key(buffer_id) {
let buf = unsafe { self.buffers.remove_and_return(buffer_id) };
let _ = self.runtime.lock().unwrap().free_buf(buf.dev_index);
}
}
pub fn host_to_pool(&mut self, src: &[u8], dst: PoolBufferId, event_wait_list: Vec<Event>) -> Result<Event, BackendError> {
let _ = event_wait_list;
let rt = &self.runtime;
let buf = self
.buffers
.get_mut(dst)
.ok_or_else(|| BackendError { status: ErrorStatus::MemoryCopyH2P, context: "invalid buffer id".into() })?;
let len = src.len().min(buf.size as usize);
let (cname, shm_ptr, _) = create_temp_shm(len as u64)?;
let shm_path = cname.to_str().unwrap_or("/none");
unsafe { std::ptr::copy_nonoverlapping(src.as_ptr(), shm_ptr, len) };
rt.lock().unwrap().write_buf(buf.dev_index, shm_path, len as u64)?;
unsafe {
libc::munmap(shm_ptr as *mut libc::c_void, len as usize);
libc::shm_unlink(cname.as_ptr());
}
Ok(Event::TT(TTEvent))
}
pub fn pool_to_host(&mut self, src: PoolBufferId, dst: &mut [u8], event_wait_list: Vec<Event>) -> Result<(), BackendError> {
let _ = event_wait_list;
let rt = &self.runtime;
let buf = self
.buffers
.get_mut(src)
.ok_or_else(|| BackendError { status: ErrorStatus::MemoryCopyP2H, context: "invalid buffer id".into() })?;
let len = dst.len().min(buf.size as usize);
let (cname, shm_ptr, _) = create_temp_shm(len as u64)?;
let shm_path = cname.to_str().unwrap_or("/none");
rt.lock().unwrap().read_buf(buf.dev_index, shm_path, len as u64)?;
unsafe {
std::ptr::copy_nonoverlapping(shm_ptr, dst.as_mut_ptr(), len);
libc::munmap(shm_ptr as *mut libc::c_void, len as usize);
libc::shm_unlink(cname.as_ptr());
}
Ok(())
}
pub fn pool_to_pool(
&mut self,
src_pool: &mut MemoryPool,
src: PoolBufferId,
dst: PoolBufferId,
event_wait_list: Vec<Event>,
) -> Result<Event, BackendError> {
match src_pool {
MemoryPool::Host(host_pool) => {
let data = host_pool.get_buffer(src);
self.host_to_pool(data, dst, event_wait_list)
}
_ => todo!(),
}
}
pub fn sync_events(&mut self, events: Vec<Event>) -> Result<(), BackendError> {
let _ = self;
let _ = events;
Ok(())
}
pub fn release_events(&mut self, events: Vec<Event>) {
let _ = self;
let _ = events;
}
pub fn dev_index(&self, buffer_id: PoolBufferId) -> Result<u32, BackendError> {
if self.buffers.contains_key(buffer_id) {
Ok(self.buffers[buffer_id].dev_index)
} else {
Err(BackendError { status: ErrorStatus::MemoryAllocation, context: "invalid buffer id".into() })
}
}
}
#[derive(Debug)]
struct RuntimeProcess {
stdin: BufWriter<ChildStdin>,
stdout: BufReader<ChildStdout>,
child: Child,
timeout_ms: u64,
}
impl RuntimeProcess {
fn new(runtime_path: &str, cache_dir: &str) -> Result<Self, BackendError> {
eprintln!("[TT_DEBUG] spawning tt-runtime from {runtime_path}");
let _ = std::process::Command::new("pkill").arg("-9").arg("zyx-tt-runtime").output();
let mut child = Command::new(runtime_path)
.stdin(std::process::Stdio::piped())
.stdout(std::process::Stdio::piped())
.stderr(std::process::Stdio::inherit())
.spawn()
.map_err(|e| BackendError {
status: ErrorStatus::Initialization,
context: format!("spawn tt-runtime {runtime_path}: {e}").into(),
})?;
eprintln!("[TT_DEBUG] child spawned, taking stdin/stdout");
let stdin = child
.stdin
.take()
.ok_or_else(|| BackendError { status: ErrorStatus::Initialization, context: "tt-runtime: no stdin".into() })?;
let stdout = child
.stdout
.take()
.ok_or_else(|| BackendError { status: ErrorStatus::Initialization, context: "tt-runtime: no stdout".into() })?;
let mut rt = RuntimeProcess { stdin: BufWriter::new(stdin), stdout: BufReader::new(stdout), child, timeout_ms: 30000 };
eprintln!("[TT_DEBUG] sending init");
let init_json = format!(r#"{{"cmd":"init","cache_dir":"{cache_dir}"}}"#);
rt.send(&init_json)?;
eprintln!("[TT_DEBUG] init sent, waiting for response");
let resp = rt.recv_with_timeout(rt.timeout_ms)?;
eprintln!("[TT_DEBUG] init response: {resp}");
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError {
status: ErrorStatus::Initialization,
context: format!("tt-runtime init error: {msg}").into(),
});
}
Ok(rt)
}
fn send(&mut self, json: &str) -> Result<(), BackendError> {
eprintln!("[RUST_SEND] {}", &json[..json.len().min(200)]);
self.stdin
.write_all(json.as_bytes())
.map_err(|e| BackendError { status: ErrorStatus::KernelLaunch, context: format!("tt-runtime write: {e}").into() })?;
self.stdin.write_all(b"\n").map_err(|e| BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime write nl: {e}").into(),
})?;
self.stdin
.flush()
.map_err(|e| BackendError { status: ErrorStatus::KernelLaunch, context: format!("tt-runtime flush: {e}").into() })?;
Ok(())
}
fn poll_read(&mut self, timeout_ms: u64) -> Result<bool, BackendError> {
match self.child.try_wait() {
Ok(Some(status)) => {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime exited unexpectedly (status {status})").into(),
});
}
Err(e) => {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime wait error: {e}").into(),
});
}
Ok(None) => {}
}
let fd = std::os::unix::io::AsRawFd::as_raw_fd(self.stdout.get_mut());
let mut pollfd = libc::pollfd { fd, events: libc::POLLIN, revents: 0 };
let timeout_ms = i32::try_from(timeout_ms).unwrap_or(i32::MAX);
let ret = unsafe { libc::poll(&mut pollfd, 1, timeout_ms) };
match ret {
-1 => {
let err = std::io::Error::last_os_error();
Err(BackendError { status: ErrorStatus::KernelLaunch, context: format!("poll error: {err}").into() })
}
0 => Ok(false),
_ => Ok(pollfd.revents & libc::POLLIN != 0),
}
}
fn recv_with_timeout(&mut self, timeout_ms: u64) -> Result<String, BackendError> {
let mut attempts = 0;
let max_attempts = 3;
let poll_timeout = timeout_ms / max_attempts;
while attempts < max_attempts {
if self.poll_read(poll_timeout)? {
let mut line = String::new();
match self.stdout.read_line(&mut line) {
Ok(0) => {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: "tt-runtime closed stdout".into(),
});
}
Ok(_) => {
let trimmed = line.trim().to_string();
if trimmed.starts_with('{') {
eprintln!("[RUST_RECV] {trimmed}");
return Ok(trimmed);
}
continue;
}
Err(_) => {
attempts += 1;
continue;
}
}
}
match self.child.try_wait() {
Ok(Some(status)) => {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime exited unexpectedly during read (status {status})").into(),
});
}
Err(e) => {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime wait error during read: {e}").into(),
});
}
Ok(None) => {
attempts += 1;
}
}
}
Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime read timeout after {}ms", timeout_ms).into(),
})
}
fn alloc_buf(&mut self, size: u64, tile_bytes: u64) -> Result<u32, BackendError> {
let cmd = format!(r#"{{"cmd":"alloc_buf","size":{size},"tile_bytes":{tile_bytes}}}"#);
self.send(&cmd)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError {
status: ErrorStatus::MemoryAllocation,
context: format!("alloc_buf error: {msg}").into(),
});
}
let idx_str = extract_json_str(&resp, "index").ok_or_else(|| BackendError {
status: ErrorStatus::MemoryAllocation,
context: "alloc_buf: no index in response".into(),
})?;
let idx: u32 = idx_str.parse().map_err(|_| BackendError {
status: ErrorStatus::MemoryAllocation,
context: format!("alloc_buf: invalid index '{idx_str}'").into(),
})?;
Ok(idx)
}
fn free_buf(&mut self, dev_index: u32) -> Result<(), BackendError> {
let cmd = format!(r#"{{"cmd":"free_buf","index":{dev_index}}}"#);
self.send(&cmd)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError { status: ErrorStatus::MemoryAllocation, context: format!("free_buf error: {msg}").into() });
}
Ok(())
}
fn write_buf(&mut self, dev_index: u32, shm_path: &str, size: u64) -> Result<(), BackendError> {
let cmd = format!(r#"{{"cmd":"write_buf","index":{dev_index},"shm_path":"{shm_path}","size":{size}}}"#);
self.send(&cmd)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError { status: ErrorStatus::MemoryCopyH2P, context: format!("write_buf error: {msg}").into() });
}
Ok(())
}
fn read_buf(&mut self, dev_index: u32, shm_path: &str, size: u64) -> Result<(), BackendError> {
let cmd = format!(r#"{{"cmd":"read_buf","index":{dev_index},"shm_path":"{shm_path}","size":{size}}}"#);
self.send(&cmd)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError { status: ErrorStatus::MemoryCopyP2H, context: format!("read_buf error: {msg}").into() });
}
Ok(())
}
fn compile_program(
&mut self,
id: u32,
reader_source: &str,
compute_source: &str,
writer_source: &str,
cb_config: &[(u32, u32, u32)],
) -> Result<(), BackendError> {
let reader_source_len = reader_source.len();
let compute_source_len = compute_source.len();
let writer_source_len = writer_source.len();
let n_cbs = cb_config.len();
let mut cmd = format!(
r#"{{"cmd":"compile_program","id":{id},"reader_source_len":{reader_source_len},"compute_source_len":{compute_source_len},"writer_source_len":{writer_source_len},"n_cbs":{n_cbs}"#
);
for (i, (idx, fmt, tb)) in cb_config.iter().enumerate() {
cmd.push_str(&format!(r#","cb_idx{i}":{idx},"cb_fmt{i}":{fmt},"cb_tb{i}":{tb}"#));
}
cmd.push('}');
self.send(&cmd)?;
self.stdin.write_all(reader_source.as_bytes()).map_err(|e| BackendError {
status: ErrorStatus::KernelCompilation,
context: format!("tt-runtime write reader: {e}").into(),
})?;
self.stdin.write_all(compute_source.as_bytes()).map_err(|e| BackendError {
status: ErrorStatus::KernelCompilation,
context: format!("tt-runtime write compute: {e}").into(),
})?;
self.stdin.write_all(writer_source.as_bytes()).map_err(|e| BackendError {
status: ErrorStatus::KernelCompilation,
context: format!("tt-runtime write writer: {e}").into(),
})?;
self.stdin.flush().map_err(|e| BackendError {
status: ErrorStatus::KernelCompilation,
context: format!("tt-runtime flush: {e}").into(),
})?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError {
status: ErrorStatus::KernelCompilation,
context: format!("tt-runtime compile error: {msg}").into(),
});
}
Ok(())
}
fn run(&mut self, id: u32, src_indices: &[u32], dst_indices: &[u32], grid_dims: [u32; 2]) -> Result<(), BackendError> {
let mut cmd = format!(r#"{{"cmd":"run","id":{id},"gd0":{gd0},"gd1":{gd1}"#, gd0 = grid_dims[0], gd1 = grid_dims[1]);
for (i, idx) in src_indices.iter().enumerate() {
cmd.push_str(&format!(r#","src{i}":{idx}"#));
}
for (i, idx) in dst_indices.iter().enumerate() {
cmd.push_str(&format!(r#","dst{i}":{idx}"#));
}
cmd.push('}');
self.send(&cmd)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime run error: {msg}").into(),
});
}
Ok(())
}
fn exit(&mut self) -> Result<(), BackendError> {
self.send(r#"{"cmd":"exit"}"#)?;
let resp = self.recv_with_timeout(self.timeout_ms)?;
if resp.contains("\"error\"") {
let msg = extract_json_str(&resp, "msg").unwrap_or_else(|| "unknown".into());
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("tt-runtime exit error: {msg}").into(),
});
}
self.child.wait().ok();
Ok(())
}
}
fn extract_json_str(json: &str, key: &str) -> Option<String> {
let k = json.find(&format!("\"{key}\""))?;
let after_colon = &json[k + key.len() + 3..]; let start = after_colon.find('"')? + 1;
let end = after_colon[start..].find('"')?;
Some(after_colon[start..start + end].to_string())
}
#[derive(Debug)]
struct TTProgram {
input_dtypes: Vec<DType>,
output_dtypes: Vec<DType>,
grid_dims: [u32; 2],
}
#[derive(Debug)]
pub struct TTDevice {
device_info: DeviceInfo,
memory_pool_id: PoolId,
runtime: Arc<Mutex<RuntimeProcess>>,
programs: Slab<DeviceProgramId, TTProgram>,
}
impl TTDevice {
pub fn deinitialize(&mut self) {}
pub const fn info(&self) -> &DeviceInfo {
&self.device_info
}
pub const fn memory_pool_id(&self) -> PoolId {
self.memory_pool_id
}
pub const fn free_compute(&self) -> u128 {
self.device_info.compute
}
#[allow(unused_must_use)]
pub fn compile(&mut self, kernel: &Kernel, debug_asm: bool) -> Result<DeviceProgramId, BackendError> {
let mut input_cb_map: Map<OpId, u32> = Map::default();
let mut output_cb_map: Map<OpId, u32> = Map::default();
let mut input_dtypes: Vec<DType> = Vec::new();
let mut output_dtypes: Vec<DType> = Vec::new();
let mut grid_dims = [1u32, 1u32];
{
let mut max_cb = 0;
let mut scan = kernel.head;
while !scan.is_null() {
match &kernel.ops[scan].op {
Op::Define { dtype, scope: Scope::Global, ro: true, .. } => input_dtypes.push(*dtype),
Op::Define { dtype, scope: Scope::Global, ro: false, .. } => output_dtypes.push(*dtype),
Op::GroupIndex { len, axis } => grid_dims[*axis as usize] = *len as u32,
Op::Store { dst, x, .. } => {
if let Op::Define { scope: Scope::Local, .. } = kernel.ops[*dst].op {
if let Op::Load { src, .. } = kernel.ops[*x].op {
if let Op::Define { ro: true, .. } = kernel.ops[src].op {
input_cb_map.insert(*dst, max_cb);
max_cb += 1;
} else {
unreachable!()
}
} else {
output_cb_map.insert(*dst, max_cb);
max_cb += 1;
}
}
}
_ => {}
}
scan = kernel.next_op(scan);
}
}
let (reader, compute, writer) =
kernel.generate_tenstorrent(debug_asm, input_dtypes.len(), output_dtypes.len(), &input_cb_map, &output_cb_map);
let prog_id = self.programs.push(TTProgram { input_dtypes, output_dtypes, grid_dims });
{
let mut cb_config = Vec::with_capacity(input_cb_map.len() + output_cb_map.len());
let dtype_to_tt_fmt = |dt: DType| -> u32 {
match dt {
DType::F32 => 0,
DType::F16 => 1,
DType::BF16 => 2,
_ => 0,
}
};
let tile_bytes_of = |dt: DType| -> u32 {
let te = 1024u64;
(match dt {
DType::F32 => 4 * te,
DType::F16 | DType::BF16 => 2 * te,
_ => 4 * te,
}) as u32
};
let mut cb_ids: Vec<u32> = input_cb_map.values().copied().collect();
for cb_id in output_cb_map.values() {
if !cb_ids.contains(cb_id) {
cb_ids.push(*cb_id);
}
}
cb_ids.sort();
for cb_id in &cb_ids {
let local_op = input_cb_map
.iter()
.find(|(_, v)| *v == cb_id)
.or_else(|| output_cb_map.iter().find(|(_, v)| *v == cb_id))
.map(|(op, _)| *op);
let dt = local_op
.and_then(|op| {
if let Op::Define { dtype, .. } = &kernel.ops[op].op {
Some(*dtype)
} else {
None
}
})
.unwrap_or(DType::BF16);
let fmt = dtype_to_tt_fmt(dt);
let tb = tile_bytes_of(dt);
cb_config.push((*cb_id, fmt, tb));
}
let mut rt_guard = self.runtime.lock().unwrap();
rt_guard.compile_program(prog_id.0, &reader, &compute, &writer, &cb_config)?;
}
Ok(prog_id)
}
pub fn release(&mut self, program_id: DeviceProgramId) {
if self.programs.contains_key(program_id) {
unsafe { self.programs.remove_and_return(program_id) };
}
}
pub fn launch(
&mut self,
program_id: DeviceProgramId,
memory_pool: &mut TTMemoryPool,
args: &[PoolBufferId],
event_wait_list: Vec<Event>,
) -> Result<Event, BackendError> {
let _ = event_wait_list;
let prog = if self.programs.contains_key(program_id) {
&self.programs[program_id]
} else {
return Err(BackendError { status: ErrorStatus::KernelLaunch, context: "invalid program id".into() });
};
let rt = &self.runtime;
let n_inputs = prog.input_dtypes.len();
let n_outputs = prog.output_dtypes.len();
if args.len() < n_inputs + n_outputs {
return Err(BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("expected {} buffers, got {}", n_inputs + n_outputs, args.len()).into(),
});
}
let mut src_indices: Vec<u32> = Vec::with_capacity(n_inputs);
for i in 0..n_inputs {
let idx = memory_pool.dev_index(args[i]).map_err(|e| BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("src{i} dev_index: {e}").into(),
})?;
src_indices.push(idx);
}
let mut dst_indices: Vec<u32> = Vec::with_capacity(n_outputs);
for i in 0..n_outputs {
let idx = memory_pool.dev_index(args[n_inputs + i]).map_err(|e| BackendError {
status: ErrorStatus::KernelLaunch,
context: format!("dst{i} dev_index: {e}").into(),
})?;
dst_indices.push(idx);
}
let mut rt_guard = rt.lock().unwrap();
rt_guard.run(program_id.0, &src_indices, &dst_indices, prog.grid_dims)?;
Ok(Event::TT(TTEvent))
}
}