I have a Tensor class:
from __future__ import annotations
from typing import Any, List, TYPE_CHECKING
import numbers
import jax
import jax.numpy as jnp
from jax.typing import DTypeLike, ArrayLike
import numpy as np
# bunch of other imports
class Tensor:
def __init__(
self,
data: ArrayLike,
name: str | None = None,
dtype: DTypeLike | None = None,
device: Any | None = None,
requires_grad: bool = False,
) -> None:
self.data: jax.Array = jax.device_put(
jnp.array(data, dtype=dtype), device=device
)
if requires_grad and not jnp.issubdtype(self.data.dtype, jnp.floating):
raise ValueError(
f"Only floating-point dtypes can have requires_grad=True. "
f"Got dtype={self.data.dtype}. Convert to float first."
)
self.name: str = name or ""
self.grads: jax.Array | None = None
self.grad_fn: GradFunction | None = None
self.requires_grad: bool = requires_grad