use std::fmt::Debug;
use std::time::Instant;
use scirs2_core::numeric::Float;
use crate::error::Result;
use super::types::{BufferFlags, BufferMetadata, DataType, MemoryLayout};
use super::DeviceId;
#[derive(Debug)]
pub struct TPUBuffer<T: Float + Debug + Send + Sync + 'static> {
pub(super) data: Vec<T>,
pub(super) shape: Vec<usize>,
layout: MemoryLayout,
device: Option<DeviceId>,
metadata: BufferMetadata,
}
impl<T: Float + Debug + Send + Sync + 'static> TPUBuffer<T> {
pub fn new(data: Vec<T>, shape: Vec<usize>, layout: MemoryLayout) -> Self {
Self {
data,
shape,
layout,
device: None,
metadata: BufferMetadata {
created_at: Instant::now(),
last_accessed: Instant::now(),
access_count: 0,
data_type: DataType::F32, flags: BufferFlags {
read_only: false,
persistent: false,
prefetch: false,
pinned: false,
},
},
}
}
pub fn layout(&self) -> MemoryLayout {
self.layout
}
pub fn size_bytes(&self) -> usize {
self.data.len() * std::mem::size_of::<T>()
}
pub fn transfer_to_device(&mut self, device: DeviceId) -> Result<()> {
self.device = Some(device);
self.metadata.last_accessed = Instant::now();
self.metadata.access_count += 1;
Ok(())
}
}