Module Nx_core.Make_frontend

Frontend functor parameterized by a backend implementation.

Parameters

module B : Backend_intf.S

Signature

module B = B
val err : string -> ('a, unit, string, 'b) Stdlib.format4 -> 'a
type ('a, 'b) t = ('a, 'b) B.t
type context = B.context
type float16_elt = Nx_buffer.float16_elt
type float32_elt = Nx_buffer.float32_elt
type float64_elt = Nx_buffer.float64_elt
type bfloat16_elt = Nx_buffer.bfloat16_elt
type float8_e4m3_elt = Nx_buffer.float8_e4m3_elt
type float8_e5m2_elt = Nx_buffer.float8_e5m2_elt
type int32_elt = Nx_buffer.int32_elt
type uint32_elt = Nx_buffer.uint32_elt
type int64_elt = Nx_buffer.int64_elt
type uint64_elt = Nx_buffer.uint64_elt
type complex32_elt = Nx_buffer.complex32_elt
type complex64_elt = Nx_buffer.complex64_elt
type bool_elt = Nx_buffer.bool_elt
type ('a, 'b) dtype = ('a, 'b) Dtype.t =
  1. | Float16 : (float, float16_elt) dtype
  2. | Float32 : (float, float32_elt) dtype
  3. | Float64 : (float, float64_elt) dtype
  4. | BFloat16 : (float, bfloat16_elt) dtype
  5. | Float8_e4m3 : (float, float8_e4m3_elt) dtype
  6. | Float8_e5m2 : (float, float8_e5m2_elt) dtype
  7. | Int4 : (int, int4_elt) dtype
  8. | UInt4 : (int, uint4_elt) dtype
  9. | Int8 : (int, int8_elt) dtype
  10. | UInt8 : (int, uint8_elt) dtype
  11. | Int16 : (int, int16_elt) dtype
  12. | UInt16 : (int, uint16_elt) dtype
  13. | Int32 : (int32, int32_elt) dtype
  14. | UInt32 : (int32, uint32_elt) dtype
  15. | Int64 : (int64, int64_elt) dtype
  16. | UInt64 : (int64, uint64_elt) dtype
  17. | Complex64 : (Stdlib.Complex.t, complex32_elt) dtype
  18. | Complex128 : (Stdlib.Complex.t, complex64_elt) dtype
  19. | Bool : (bool, bool_elt) dtype
type float16_t = (float, float16_elt) t
type float32_t = (float, float32_elt) t
type float64_t = (float, float64_elt) t
type int8_t = (int, int8_elt) t
type uint8_t = (int, uint8_elt) t
type int16_t = (int, int16_elt) t
type uint16_t = (int, uint16_elt) t
type int32_t = (int32, int32_elt) t
type int64_t = (int64, int64_elt) t
type uint32_t = (int32, uint32_elt) t
type uint64_t = (int64, uint64_elt) t
type complex64_t = (Stdlib.Complex.t, complex32_elt) t
type complex128_t = (Stdlib.Complex.t, complex64_elt) t
type bool_t = (bool, bool_elt) t
val float16 : (float, float16_elt) dtype
val float32 : (float, float32_elt) dtype
val float64 : (float, float64_elt) dtype
val bfloat16 : (float, bfloat16_elt) dtype
val float8_e4m3 : (float, float8_e4m3_elt) dtype
val float8_e5m2 : (float, float8_e5m2_elt) dtype
val int4 : (int, int4_elt) dtype
val uint4 : (int, uint4_elt) dtype
val int8 : (int, int8_elt) dtype
val uint8 : (int, uint8_elt) dtype
val int16 : (int, int16_elt) dtype
val uint16 : (int, uint16_elt) dtype
val int32 : (int32, int32_elt) dtype
val uint32 : (int32, uint32_elt) dtype
val int64 : (int64, int64_elt) dtype
val uint64 : (int64, uint64_elt) dtype
val complex64 : (Stdlib.Complex.t, complex32_elt) dtype
val complex128 : (Stdlib.Complex.t, complex64_elt) dtype
val bool : (bool, bool_elt) dtype
type index =
  1. | I of int
  2. | L of int list
  3. | R of int * int
  4. | Rs of int * int * int
  5. | A
  6. | M of (bool, bool_elt) t
  7. | N
val data : ('a, 'b) B.t -> ('a, 'b) Nx_buffer.t
val shape : ('a, 'b) B.t -> int array
val dtype : ('a, 'b) B.t -> ('a, 'b) Dtype.t
val itemsize : ('a, 'b) B.t -> int
val strides : ('a, 'b) B.t -> int array
val stride : int -> ('a, 'b) B.t -> int
val dims : ('a, 'b) B.t -> int array
val dim : int -> ('a, 'b) B.t -> int
val ndim : ('a, 'b) B.t -> int
val size : ('a, 'b) B.t -> int
val numel : ('a, 'b) B.t -> int
val nbytes : ('a, 'b) B.t -> int
val offset : ('a, 'b) B.t -> int
val is_c_contiguous : ('a, 'b) B.t -> bool
val array_prod : int array -> int
module IntSet : sig ... end
val power_of_two : 'a 'b. ('a, 'b) Dtype.t -> int -> 'a
val ensure_float_dtype : string -> ('a, 'b) B.t -> unit
val ensure_int_dtype : string -> ('a, 'b) B.t -> unit
val resolve_axis : ?ndim_opt:??? -> ('a, 'b) B.t -> int option -> int array
val resolve_single_axis : ?ndim_opt:??? -> ('a, 'b) B.t -> int -> int
val normalize_and_dedup_axes : op:string -> int -> int list -> int list
val reduction_element_count : int array -> ?axes:??? -> unit -> int
val copy_to_out : 'a -> 'a
val reshape : Shape.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val broadcast_shapes : Shape.t -> Shape.t -> int array
val broadcast_to : Shape.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val broadcasted : ?reverse:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t * ('a, 'b) B.t
val expand : int array -> ('a, 'b) B.t -> ('a, 'b) B.t
val cast : ('c, 'd) Dtype.t -> ('a, 'b) t -> ('c, 'd) t
val astype : ('a, 'b) Dtype.t -> ('c, 'd) t -> ('a, 'b) t
val contiguous : ('a, 'b) B.t -> ('a, 'b) B.t
val copy : ('a, 'b) B.t -> ('a, 'b) B.t
val blit : ('a, 'b) B.t -> ('a, 'b) B.t -> unit
val create : B.context -> ('a, 'b) Dtype.t -> int array -> 'a array -> ('a, 'b) B.t
val init : B.context -> ('a, 'b) Dtype.t -> Shape.t -> (int array -> 'a) -> ('a, 'b) B.t
val scalar : B.context -> ('a, 'b) Dtype.t -> 'a -> ('a, 'b) B.t
val scalar_like : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val fill : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val empty : B.context -> ('a, 'b) Dtype.t -> int array -> ('a, 'b) B.t
val zeros : B.context -> ('a, 'b) Dtype.t -> int array -> ('a, 'b) B.t
val ones : B.context -> ('a, 'b) Dtype.t -> int array -> ('a, 'b) B.t
val full : B.context -> ('a, 'b) Dtype.t -> int array -> 'a -> ('a, 'b) B.t
val create_like : ('a, 'b) B.t -> (B.context -> ('a, 'b) Dtype.t -> int array -> 'c) -> 'c
val empty_like : ('a, 'b) B.t -> ('a, 'b) B.t
val full_like : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val zeros_like : ('a, 'b) B.t -> ('a, 'b) B.t
val ones_like : ('a, 'b) B.t -> ('a, 'b) B.t
val to_buffer : ('a, 'b) B.t -> ('a, 'b) Nx_buffer.t
val to_bigarray : ('c, 'd) B.t -> ('a, 'b, Stdlib.Bigarray.c_layout) Stdlib.Bigarray.Genarray.t
val of_buffer : B.context -> shape:Shape.t -> ('a, 'b) Nx_buffer.t -> ('a, 'b) B.t
val of_bigarray : B.context -> 'c -> ('a, 'b) B.t
val to_array : ('a, 'b) B.t -> 'a array
val binop : (('a, 'b) B.t -> ('a, 'b) B.t -> 'c) -> ('a, 'b) B.t -> ('a, 'b) B.t -> 'c
val cmpop : (('a, 'b) B.t -> ('a, 'b) B.t -> 'c) -> ('a, 'b) B.t -> ('a, 'b) B.t -> 'c
val add : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val add_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val radd_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val sub : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val sub_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rsub_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val mul : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val mul_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rmul_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val div : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val div_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rdiv_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val pow : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val pow_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rpow_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val maximum : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val maximum_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rmaximum_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val minimum : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val minimum_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rminimum_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val mod_ : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val mod_s : ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val rmod_s : 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val bitwise_xor : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val bitwise_or : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val bitwise_and : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val logical_and : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val logical_or : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val logical_xor : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val logical_not : ('a, 'b) t -> ('a, 'b) t
val cmpeq : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val cmpne : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val cmplt : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val cmple : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val cmpgt : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val cmpge : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val less : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val less_equal : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val greater : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val greater_equal : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val equal : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val not_equal : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val equal_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val not_equal_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val less_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val greater_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val less_equal_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val greater_equal_s : ('a, 'b) B.t -> 'a -> (bool, Dtype.bool_elt) B.t
val unaryop : ('a -> 'b) -> 'a -> 'b
val neg : ('a, 'b) B.t -> ('a, 'b) B.t
val bitwise_not : ('a, 'b) B.t -> ('a, 'b) B.t
val invert : ('a, 'b) B.t -> ('a, 'b) B.t
val sin : ('a, 'b) B.t -> ('a, 'b) B.t
val cos : ('a, 'b) B.t -> ('a, 'b) B.t
val sqrt : ('a, 'b) B.t -> ('a, 'b) B.t
val recip : ('a, 'b) B.t -> ('a, 'b) B.t
val log : ('a, 'b) B.t -> ('a, 'b) B.t
val exp : ('a, 'b) B.t -> ('a, 'b) B.t
val abs : ('a, 'b) B.t -> ('a, 'b) B.t
val log2 : ('a, 'b) B.t -> ('a, 'b) B.t
val exp2 : ('a, 'b) B.t -> ('a, 'b) B.t
val tan : ('a, 'b) B.t -> ('a, 'b) B.t
val square : ('a, 'b) B.t -> ('a, 'b) B.t
val sign : ('a, 'b) B.t -> ('a, 'b) B.t
val relu : ('a, 'b) B.t -> ('a, 'b) B.t
val sigmoid : ('a, 'b) B.t -> ('a, 'b) B.t
val rsqrt : ('a, 'b) B.t -> ('a, 'b) B.t
val asin : ('a, 'b) B.t -> ('a, 'b) B.t
val acos : ('a, 'b) B.t -> ('a, 'b) B.t
val atan : ('a, 'b) B.t -> ('a, 'b) B.t
val sinh : ('a, 'b) B.t -> ('a, 'b) B.t
val cosh : ('a, 'b) B.t -> ('a, 'b) B.t
val tanh : ('a, 'b) B.t -> ('a, 'b) B.t
val asinh : ('a, 'b) B.t -> ('a, 'b) B.t
val acosh : ('a, 'b) B.t -> ('a, 'b) B.t
val atanh : ('a, 'b) B.t -> ('a, 'b) B.t
val trunc : ('a, 'b) B.t -> ('a, 'b) B.t
val ceil : ('a, 'b) B.t -> ('a, 'b) B.t
val floor : ('a, 'b) B.t -> ('a, 'b) B.t
val round : ('a, 'b) B.t -> ('a, 'b) B.t
val isinf : ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val isnan : ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val isfinite : ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val lerp : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val lerp_scalar_weight : ('a, 'b) B.t -> ('a, 'b) B.t -> 'a -> ('a, 'b) B.t
val shift_op : op:string -> apply:(('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t) -> ('a, 'b) B.t -> int -> ('a, 'b) B.t
val lshift : ('a, 'b) B.t -> int -> ('a, 'b) B.t
val rshift : ('a, 'b) B.t -> int -> ('a, 'b) B.t
val clamp : ?min:??? -> ?max:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val clip : ?min:??? -> ?max:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val where : (bool, Dtype.bool_elt) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val atan2 : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val hypot : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val reduce_op : (axes:int array -> keepdims:bool -> ('a, 'b) B.t -> 'c) -> ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> 'c
val sum : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val max : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val min : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val prod : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val associative_scan : axis:int -> [ `Max | `Min | `Prod | `Sum ] -> ('a, 'b) B.t -> ('a, 'b) B.t
val cumulative_scan : ?axis:??? -> [ `Max | `Min | `Prod | `Sum ] -> ('a, 'b) B.t -> ('a, 'b) B.t
val cumsum : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val cumprod : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val cummax : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val cummin : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val mean : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val var : ?axes:??? -> ?keepdims:??? -> ?ddof:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val std : ?axes:??? -> ?keepdims:??? -> ?ddof:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val all : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val any : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val array_equal : ('a, 'b) B.t -> ('a, 'b) B.t -> (bool, Dtype.bool_elt) B.t
val pad : (int * int) array -> 'a -> ('a, 'b) B.t -> ('a, 'b) B.t
val shrink : (int * int) array -> ('a, 'b) B.t -> ('a, 'b) B.t
val flatten : ?start_dim:??? -> ?end_dim:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val unflatten : int -> int array -> ('a, 'b) B.t -> ('a, 'b) B.t
val ravel : ('a, 'b) B.t -> ('a, 'b) B.t
val squeeze : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val unsqueeze : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val squeeze_axis : IntSet.elt -> ('a, 'b) B.t -> ('a, 'b) B.t
val unsqueeze_axis : IntSet.elt -> ('a, 'b) B.t -> ('a, 'b) B.t
val expand_dims : IntSet.elt list -> ('a, 'b) B.t -> ('a, 'b) B.t
val transpose : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val flip : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val moveaxis : int -> int -> ('a, 'b) B.t -> ('a, 'b) B.t
val swapaxes : int -> int -> ('a, 'b) B.t -> ('a, 'b) B.t
val cat_tensors : axis:int -> ('a, 'b) B.t list -> ('a, 'b) B.t
val roll : ?axis:??? -> int -> ('a, 'b) B.t -> ('a, 'b) B.t
val tile : int array -> ('a, 'b) B.t -> ('a, 'b) B.t
val repeat : ?axis:??? -> int -> ('a, 'b) B.t -> ('a, 'b) B.t
val check_dtypes_match : op:string -> ('a, 'b) B.t list -> unit
val concatenate : ?axis:??? -> ('a, 'b) B.t list -> ('a, 'b) B.t
val stack : ?axis:??? -> ('a, 'b) B.t list -> ('a, 'b) B.t
val ensure_ndim : int -> ('a, 'b) B.t -> ('a, 'b) B.t
val vstack : ('a, 'b) B.t list -> ('a, 'b) B.t
val hstack : ('a, 'b) B.t list -> ('a, 'b) B.t
val dstack : ('a, 'b) B.t list -> ('a, 'b) B.t
val broadcast_arrays : ('a, 'b) B.t list -> ('a, 'b) B.t list
val eye : B.context -> ?m:??? -> ?k:??? -> ('a, 'b) Dtype.t -> int -> ('a, 'b) B.t
val identity : B.context -> ('a, 'b) Dtype.t -> int -> ('a, 'b) B.t
val diag : ?k:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val arange : B.context -> ('a, 'b) Dtype.t -> int -> int -> int -> ('a, 'b) B.t
val arange_f : B.context -> (float, 'a) Dtype.t -> float -> float -> float -> (float, 'a) B.t
val linspace : B.context -> ('a, 'b) Dtype.t -> ?endpoint:??? -> float -> float -> int -> ('a, 'b) B.t
val logspace : B.context -> (float, 'a) Dtype.t -> ?endpoint:??? -> ?base:??? -> float -> float -> int -> (float, 'a) B.t
val geomspace : B.context -> (float, 'a) Dtype.t -> ?endpoint:??? -> float -> float -> int -> (float, 'a) B.t
val meshgrid : ?indexing:??? -> ('a, 'b) B.t -> ('c, 'd) B.t -> ('a, 'b) B.t * ('c, 'd) B.t
val triangular_mask : op:string -> cmp: ((int32, int32_elt) B.t -> (int32, int32_elt) B.t -> (bool, Dtype.bool_elt) B.t) -> ?k:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val tril : ?k:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val triu : ?k:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val apply_index_mode : mode:[< `clip | `raise | `wrap ] -> n:int -> B.context -> (int32, int32_elt) B.t -> (int32, int32_elt) B.t
val take : ?axis:??? -> ?mode:??? -> (int32, int32_elt) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val take_along_axis : axis:int -> (int32, Dtype.int32_elt) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val normalize_index : int -> int -> int
val normalize_and_check_index : op:string -> int -> int -> int
type dim_op =
  1. | View of {
    1. start : int;
    2. stop : int;
    3. step : int;
    4. dim_len : int;
    }
  2. | Squeeze of {
    1. idx : int;
    }
  3. | Gather of int array
  4. | New_axis
val normalize_slice_spec : int -> index -> dim_op
val slice_internal : index list -> ('a, 'b) B.t -> ('a, 'b) B.t
val set_slice_internal : index list -> ('a, 'b) B.t -> ('a, 'b) B.t -> unit
val get : int list -> ('a, 'b) B.t -> ('a, 'b) B.t
val set : int list -> ('a, 'b) B.t -> ('a, 'b) B.t -> unit
val unsafe_get : int list -> ('a, 'b) B.t -> 'a
val unsafe_set : int list -> 'a -> ('a, 'b) B.t -> unit
val slice : index list -> ('a, 'b) B.t -> ('a, 'b) B.t
val set_slice : index list -> ('a, 'b) B.t -> ('a, 'b) B.t -> unit
val item : int list -> ('a, 'b) B.t -> 'a
val set_item : int list -> 'a -> ('a, 'b) B.t -> unit
val put : ?axis:??? -> indices:(int32, int32_elt) B.t -> values:('a, 'b) B.t -> ?mode:??? -> ('a, 'b) B.t -> unit
val index_put : indices:(int32, int32_elt) B.t array -> values:('a, 'b) B.t -> ?mode:??? -> ('a, 'b) B.t -> unit
val scatter : ?mode:??? -> ?unique_indices:??? -> axis:int -> indices:(int32, Dtype.int32_elt) B.t -> values:('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val put_along_axis : axis:int -> indices:(int32, Dtype.int32_elt) B.t -> values:('a, 'b) B.t -> ('a, 'b) B.t -> unit
val nonzero_indices_only : (bool, bool_elt) t -> (int32, int32_elt) B.t array
val compress : ?axis:??? -> condition:(bool, bool_elt) t -> ('a, 'b) B.t -> ('a, 'b) B.t
val extract : condition:(bool, bool_elt) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val nonzero : ('a, 'b) t -> (int32, int32_elt) B.t array
val argwhere : ('a, 'b) t -> (int32, int32_elt) B.t
val array_split : axis:int -> [< `Count of int | `Indices of int list ] -> ('a, 'b) B.t -> ('a, 'b) B.t list
val split : axis:int -> int -> ('a, 'b) B.t -> ('a, 'b) B.t list
val sort : ?descending:??? -> ?axis:??? -> ('a, 'b) t -> ('a, 'b) t * (int32, Dtype.int32_elt) B.t
val argsort : ?descending:??? -> ?axis:??? -> ('a, 'b) t -> (int32, Dtype.int32_elt) B.t
val argmax : ?axis:??? -> ?keepdims:??? -> ('a, 'b) B.t -> (int32, Dtype.int32_elt) B.t
val argmin : ?axis:??? -> ?keepdims:??? -> ('a, 'b) t -> (int32, Dtype.int32_elt) t
val validate_random_float_params : string -> ('a, 'b) Dtype.t -> Shape.t -> unit
val rand : B.context -> ('a, 'b) Dtype.t -> Shape.t -> ('a, 'b) B.t
val randn : B.context -> ('a, 'b) Dtype.t -> Shape.t -> ('a, 'b) B.t
val randint : B.context -> ('a, 'b) Dtype.t -> ?high:??? -> Shape.t -> int -> ('a, 'b) t
val bernoulli : B.context -> p:float -> Shape.t -> (bool, Dtype.bool_elt) B.t
val permutation : B.context -> int -> (int32, Dtype.int32_elt) B.t
val shuffle : B.context -> ('a, 'b) B.t -> ('a, 'b) B.t
val categorical : B.context -> ?axis:??? -> ?shape:??? -> ('a, 'b) t -> (int32, Dtype.int32_elt) t
val truncated_normal : B.context -> ('a, 'b) Dtype.t -> lower:float -> upper:float -> Shape.t -> ('a, 'b) B.t
val matmul_with_alloc : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val dot : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val matmul : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val diagonal : ?offset:??? -> ?axis1:??? -> ?axis2:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val matrix_transpose : ('a, 'b) B.t -> ('a, 'b) B.t
val extract_complex_part : op:string -> field:(Stdlib.Complex.t -> float) -> ('a, 'b) t -> ('c, 'd) t
val complex : real:('a, 'b) t -> imag:('a, 'b) t -> 'c
val real : ('a, 'b) t -> ('c, 'd) t
val imag : ('a, 'b) t -> ('c, 'd) t
val conjugate : ('a, 'b) t -> ('a, 'b) t
val vdot : ('a, 'b) t -> ('a, 'b) t -> ('a, 'b) B.t
val vecdot : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val inner : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val outer : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val tensordot : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
module Einsum : sig ... end
val einsum : string -> ('a, 'b) B.t array -> ('a, 'b) B.t
val kron : ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val multi_dot : ('a, 'b) B.t array -> ('a, 'b) B.t
val cross : ?axis:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val check_square : op:string -> ('a, 'b) B.t -> unit
val check_float_or_complex : op:string -> ('a, 'b) t -> unit
val check_real : op:string -> ('a, 'b) t -> unit
val cholesky : ?upper:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val qr : ?mode:??? -> ('a, 'b) t -> ('a, 'b) B.t * ('a, 'b) B.t
val svd : ?full_matrices:??? -> ('a, 'b) t -> ('a, 'b) B.t * (float, Dtype.float64_elt) B.t * ('a, 'b) B.t
val svdvals : ('a, 'b) t -> (float, Dtype.float64_elt) B.t
val eig : ('a, 'b) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) B.t * (Stdlib.Complex.t, Dtype.complex64_elt) B.t
val eigh : ?uplo:??? -> ('b, 'c) B.t -> (float, Dtype.float64_elt) B.t * ('b, 'c) B.t
val eigvals : ('a, 'b) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) B.t
val eigvalsh : ?uplo:??? -> ('b, 'c) B.t -> (float, Dtype.float64_elt) B.t
val norm : ?ord:??? -> ?axes:??? -> ?keepdims:??? -> ('a, 'b) t -> ('a, 'b) B.t
val slogdet : ('a, 'b) B.t -> (float, Dtype.float32_elt) B.t * (float, Dtype.float32_elt) t
val det : ('a, 'b) B.t -> ('a, 'b) B.t
val matrix_rank : ?tol:??? -> ?rtol:??? -> ?hermitian:??? -> ('a, 'b) t -> int
val trace : ?offset:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val solve : ('a, 'b) B.t -> ('a, 'b) t -> ('a, 'b) B.t
val pinv : ?rtol:??? -> ?hermitian:??? -> ('a, 'b) t -> ('a, 'b) B.t
val lstsq : ?rcond:??? -> ('a, 'b) t -> ('a, 'b) t -> ('a, 'b) B.t * ('a, 'b) B.t * int * (float, Dtype.float64_elt) B.t
val inv : ('a, 'b) B.t -> ('a, 'b) B.t
val matrix_power : ('a, 'b) B.t -> int -> ('a, 'b) B.t
val cond : ?p:??? -> ('a, 'b) B.t -> ('a, 'b) t
val tensorsolve : ?axes:??? -> ('a, 'b) t -> ('a, 'b) t -> ('a, 'b) B.t
val tensorinv : ?ind:??? -> ('a, 'b) t -> ('a, 'b) B.t
type fft_norm = [
  1. | `Backward
  2. | `Forward
  3. | `Ortho
]
val pad_or_truncate_for_fft : ('a, 'b) B.t -> int list -> int list option -> ('a, 'b) B.t
val fft_norm_scale : [< `Backward | `Forward | `Ortho ] -> int list -> ('a, 'b) B.t -> float
val ifft_norm_scale : [< `Backward | `Forward | `Ortho ] -> int list -> ('a, 'b) B.t -> float
val apply_fft_scale : float -> (Stdlib.Complex.t, 'a) t -> (Stdlib.Complex.t, 'a) t
val fft : ?axis:??? -> ?n:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) t -> (Stdlib.Complex.t, 'a) t
val ifft : ?axis:??? -> ?n:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) t -> (Stdlib.Complex.t, 'a) t
val rfft : ?axis:??? -> ?n:??? -> ?norm:??? -> (float, 'a) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) t
val irfft : ?axis:??? -> ?n:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (float, Dtype.float64_elt) B.t
val check_fft2 : op:string -> ('a, 'b) B.t -> int list option -> int list
val fft2 : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (Stdlib.Complex.t, 'a) t
val ifft2 : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (Stdlib.Complex.t, 'a) t
val fftn : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (Stdlib.Complex.t, 'a) t
val ifftn : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (Stdlib.Complex.t, 'a) t
val rfft2 : ?axes:??? -> ?s:??? -> ?norm:??? -> (float, 'a) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) t
val irfft2 : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (float, Dtype.float64_elt) B.t
val rfftn : ?axes:??? -> ?s:??? -> ?norm:??? -> (float, 'a) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) t
val irfftn : ?axes:??? -> ?s:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (float, Dtype.float64_elt) B.t
val hfft : ?axis:??? -> ?n:??? -> ?norm:??? -> (Stdlib.Complex.t, 'a) B.t -> (float, Dtype.float64_elt) B.t
val ihfft : ?axis:??? -> ?n:??? -> ?norm:??? -> (float, 'a) B.t -> (Stdlib.Complex.t, Dtype.complex64_elt) t
val fftfreq : B.context -> ?d:??? -> int -> (float, Dtype.float64_elt) B.t
val rfftfreq : B.context -> ?d:??? -> int -> (float, Dtype.float64_elt) B.t
val fftshift : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val ifftshift : ?axes:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val softmax : ?axes:??? -> ?scale:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val log_softmax : ?axes:??? -> ?scale:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val logsumexp : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val logmeanexp : ?axes:??? -> ?keepdims:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val standardize : ?axes:??? -> ?mean:??? -> ?variance:??? -> ?epsilon:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val erf : ('a, 'b) B.t -> ('a, 'b) B.t
val extract_patches : kernel_size:int array -> stride:int array -> dilation:int array -> padding:(int * int) array -> ('a, 'b) B.t -> ('a, 'b) B.t
val combine_patches : output_size:int array -> kernel_size:int array -> stride:int array -> dilation:int array -> padding:(int * int) array -> ('a, 'b) B.t -> ('a, 'b) B.t
val correlate_padding : mode:[< `Full | `Same | `Valid ] -> 'a -> int array -> (int * int) array
val correlate : ?padding:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val convolve : ?padding:??? -> ('a, 'b) B.t -> ('a, 'b) B.t -> ('a, 'b) B.t
val sliding_filter : reduce_fn:(('a, 'b) B.t -> axes:int list -> keepdims:bool -> ('c, 'd) B.t) -> kernel_size:int array -> ?stride:??? -> ('a, 'b) B.t -> ('c, 'd) B.t
val maximum_filter : kernel_size:int array -> ?stride:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val minimum_filter : kernel_size:int array -> ?stride:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val uniform_filter : kernel_size:int array -> ?stride:??? -> ('a, 'b) B.t -> ('a, 'b) B.t
val one_hot : num_classes:int -> ('a, 'b) B.t -> (int, Dtype.uint8_elt) t
val pp_data : Stdlib.Format.formatter -> ('a, 'b) t -> unit
val format_to_string : (Stdlib.Format.formatter -> 'a -> 'b) -> 'a -> string
val print_with_formatter : (Stdlib.Format.formatter -> 'a -> 'b) -> 'a -> unit
val data_to_string : ('a, 'b) t -> string
val print_data : ('a, 'b) t -> unit
val pp_dtype : Stdlib.Format.formatter -> ('a, 'b) Dtype.t -> unit
val dtype_to_string : ('a, 'b) Dtype.t -> string
val shape_to_string : int array -> string
val pp_shape : Stdlib.Format.formatter -> int array -> unit
val pp : Stdlib.Format.formatter -> ('a, 'b) B.t -> unit
val print : ('a, 'b) B.t -> unit
val to_string : ('a, 'b) B.t -> string
val map_item : ('a -> 'a) -> ('a, 'b) B.t -> ('a, 'b) B.t
val iter_item : ('a -> 'b) -> ('a, 'c) B.t -> unit
val fold_item : ('a -> 'b -> 'a) -> 'a -> ('b, 'c) B.t -> 'a
val map : (('a, 'b) B.t -> ('a, 'b) B.t) -> ('a, 'b) B.t -> ('a, 'b) B.t
val iter : (('a, 'b) B.t -> 'c) -> ('a, 'b) B.t -> unit
val fold : ('a -> ('b, 'c) B.t -> 'a) -> 'a -> ('b, 'c) B.t -> 'a
module Infix : sig ... end