use alloc::vec::Vec;
use burn_backend::{TensorMetadata, TensorPrimitive, get_device_settings};
use burn_dispatch::{Dispatch, DispatchTensor};
use burn_std::DeviceSettings;
#[derive(Clone, Debug)]
pub struct Float;
#[derive(Clone, Debug)]
pub struct Int;
#[derive(Clone, Debug)]
pub struct Bool;
mod sealed {
pub trait Sealed {}
}
impl sealed::Sealed for Float {}
impl sealed::Sealed for Int {}
impl sealed::Sealed for Bool {}
pub trait TensorKind: sealed::Sealed + Clone + Send + Sync + core::fmt::Debug {
const KIND: Kind;
fn name() -> &'static str {
Self::KIND.as_str()
}
}
impl TensorKind for Float {
const KIND: Kind = Kind::Float;
}
impl TensorKind for Int {
const KIND: Kind = Kind::Int;
}
impl TensorKind for Bool {
const KIND: Kind = Kind::Bool;
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Kind {
Float,
Int,
Bool,
}
impl Kind {
pub fn as_str(&self) -> &'static str {
match self {
Kind::Float => "Float",
Kind::Int => "Int",
Kind::Bool => "Bool",
}
}
}
pub struct BridgeTensor {
blob: bridge_opaque::Opaque,
}
burn_std::obfuscate!(
type: BridgeTensorVariant,
module: bridge_opaque,
derives: [Send, Sync]
);
impl core::fmt::Debug for BridgeTensor {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("BridgeTensor")
.field("kind", &self.kind())
.finish()
}
}
impl Clone for BridgeTensor {
fn clone(&self) -> Self {
Self::new(self.as_variant().clone())
}
}
impl BridgeTensor {
fn as_variant(&self) -> &BridgeTensorVariant {
self.blob.as_ref()
}
fn into_variant(self) -> BridgeTensorVariant {
self.blob.into_inner()
}
fn new(inner: BridgeTensorVariant) -> Self {
Self {
blob: bridge_opaque::Opaque::new(inner),
}
}
}
#[derive(Clone, Debug)]
enum BridgeTensorVariant {
Bool(DispatchTensor),
Int(DispatchTensor),
Float(DispatchTensor),
QFloat(DispatchTensor),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BridgeKind {
Bool,
Int,
Float,
QFloat,
}
macro_rules! ext_fn {
($(#[$meta:meta])* fn $($tt:tt)+) => {
#[cfg(feature = "extension")]
$(#[$meta])*
pub fn $($tt)+
#[cfg(not(feature = "extension"))]
$(#[$meta])*
pub(crate) fn $($tt)+
};
}
impl BridgeTensor {
ext_fn! {
fn float(tensor: DispatchTensor) -> Self {
Self::new(BridgeTensorVariant::Float(tensor))
}
}
ext_fn! {
fn int(tensor: DispatchTensor) -> Self {
Self::new(BridgeTensorVariant::Int(tensor))
}
}
ext_fn! {
fn bool(tensor: DispatchTensor) -> Self {
Self::new(BridgeTensorVariant::Bool(tensor))
}
}
ext_fn! {
fn qfloat(tensor: DispatchTensor) -> Self {
Self::new(BridgeTensorVariant::QFloat(tensor))
}
}
pub fn kind(&self) -> BridgeKind {
match self.as_variant() {
BridgeTensorVariant::Bool(_) => BridgeKind::Bool,
BridgeTensorVariant::Int(_) => BridgeKind::Int,
BridgeTensorVariant::Float(_) => BridgeKind::Float,
BridgeTensorVariant::QFloat(_) => BridgeKind::QFloat,
}
}
pub fn is_float(&self) -> bool {
matches!(self.kind(), BridgeKind::Float)
}
pub fn is_int(&self) -> bool {
matches!(self.kind(), BridgeKind::Int)
}
pub fn is_bool(&self) -> bool {
matches!(self.kind(), BridgeKind::Bool)
}
pub fn is_qfloat(&self) -> bool {
matches!(self.kind(), BridgeKind::QFloat)
}
ext_fn! {
fn into_parts(self) -> (BridgeKind, DispatchTensor) {
match self.into_variant() {
BridgeTensorVariant::Bool(t) => (BridgeKind::Bool, t),
BridgeTensorVariant::Int(t) => (BridgeKind::Int, t),
BridgeTensorVariant::Float(t) => (BridgeKind::Float, t),
BridgeTensorVariant::QFloat(t) => (BridgeKind::QFloat, t),
}
}
}
ext_fn! {
fn as_parts(&self) -> (BridgeKind, &DispatchTensor) {
match self.as_variant() {
BridgeTensorVariant::Bool(t) => (BridgeKind::Bool, t),
BridgeTensorVariant::Int(t) => (BridgeKind::Int, t),
BridgeTensorVariant::Float(t) => (BridgeKind::Float, t),
BridgeTensorVariant::QFloat(t) => (BridgeKind::QFloat, t),
}
}
}
pub fn dtype(&self) -> burn_std::DType {
match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor.dtype(),
BridgeTensorVariant::Int(tensor) => tensor.dtype(),
BridgeTensorVariant::Float(tensor) => tensor.dtype(),
BridgeTensorVariant::QFloat(tensor) => tensor.dtype(),
}
}
pub fn can_mut(&self) -> bool {
match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor.can_mut(),
BridgeTensorVariant::Int(tensor) => tensor.can_mut(),
BridgeTensorVariant::Float(tensor) => tensor.can_mut(),
BridgeTensorVariant::QFloat(tensor) => tensor.can_mut(),
}
}
pub fn shape(&self) -> burn_std::Shape {
match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor.shape(),
BridgeTensorVariant::Int(tensor) => tensor.shape(),
BridgeTensorVariant::Float(tensor) => tensor.shape(),
BridgeTensorVariant::QFloat(tensor) => tensor.shape(),
}
}
pub fn rank(&self) -> usize {
match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor.rank(),
BridgeTensorVariant::Int(tensor) => tensor.rank(),
BridgeTensorVariant::Float(tensor) => tensor.rank(),
BridgeTensorVariant::QFloat(tensor) => tensor.rank(),
}
}
pub(crate) fn as_dispatch(&self) -> &DispatchTensor {
match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor,
BridgeTensorVariant::Int(tensor) => tensor,
BridgeTensorVariant::Float(tensor) => tensor,
BridgeTensorVariant::QFloat(tensor) => tensor,
}
}
#[cfg(feature = "autodiff")]
pub(crate) fn as_float(&self) -> &DispatchTensor {
match self.as_variant() {
BridgeTensorVariant::Float(tensor) => tensor,
_ => panic!("Should be Float primitive kind"),
}
}
pub(crate) fn into_dispatch_vec(tensors: Vec<Self>) -> Vec<DispatchTensor> {
tensors.into_iter().map(Into::into).collect()
}
pub(crate) fn into_float(self) -> DispatchTensor {
match self.into_variant() {
BridgeTensorVariant::Float(tensor) => tensor,
BridgeTensorVariant::QFloat(tensor) => {
TensorPrimitive::<Dispatch>::QFloat(tensor).tensor()
}
_ => panic!("Should be Float primitive kind"),
}
}
pub(crate) fn device_settings(&self) -> DeviceSettings {
let device = match self.as_variant() {
BridgeTensorVariant::Bool(tensor) => tensor.device(),
BridgeTensorVariant::Int(tensor) => tensor.device(),
BridgeTensorVariant::Float(tensor) => tensor.device(),
BridgeTensorVariant::QFloat(tensor) => tensor.device(),
};
get_device_settings::<Dispatch>(&device)
}
}
impl From<BridgeTensor> for DispatchTensor {
fn from(value: BridgeTensor) -> Self {
match value.into_variant() {
BridgeTensorVariant::Bool(tensor) => tensor,
BridgeTensorVariant::Int(tensor) => tensor,
BridgeTensorVariant::Float(tensor) => tensor,
BridgeTensorVariant::QFloat(tensor) => tensor,
}
}
}