Repository navigation
[RFC]: Updates for float16/bfloat16 and for dtypes that are lacking full support in libraries #998
Description
Activity
Let me also link the canonical issue for
bfloat16support in NumPy: numpy/numpy#19808One note: complex32/bcomplex32 has been merged into ml_dtypes, but is not yet part of a release (partly due to questions about dtype character codes). When it is released, we will integrate support into JAX.
Reacted by Ralf Gommers, Leo Fang and Athan- addedAPI changeChanges to existing functions or objects in the API.Changes to existing functions or objects in the API.
on Mar 5, 2026 - changed the title
[-]Updates for float16/bfloat16 and for dtypes that are lacking full support in libraries[/-][+][RFC]: Updates for float16/bfloat16 and for dtypes that are lacking full support in libraries[/+]on Mar 5, 2026 - addedRFCRequest for comments. Feature requests and proposed changes.Request for comments. Feature requests and proposed changes.
on Mar 5, 2026 After starting a prototype support in data-apis/array-api-strict#206, ISTM that the list of supported dtypes should be device-specific: we need some wording to state that a conforming library can have devices which only support a subset of dtypes.
The inspection API already supports this (___array_namespace_info__().dtypesaccepts thedeviceargument), what's needed additionally is for creation functions to:- state that using disallowed dtype/device combinations raises an error:
torch.zeros(3, dtype=torch.float64, device="mps")raises a TypeError; the standard may either follow it, or mandate a ValueError or keep the precise error class unspecified; - state that
dtype=Noneconstructs an array of the default dtype for the device:xp.zeros(3, device=X)returns an array of default floats for thedevice=X. It currently only says the default floating-point data type.
For array indexing, it would be helpful to clarify if:
__getitem__with the indexer array's device differing from the indexed array device: allowed, prohibited or unspecified? Data point:pytorchallows any combinations, the result's device is the device of the indexed array__setitem__: if r.h.s. device differs from the l.h.s. device, is it an error or an implicit device transfer? Torch again allows the latter, cf Raise on implicit device transfer in__setitem__array-api-strict#207; currently the spec makes a "recommendation" to not allow implicit transfers, Raise on implicit device transfer in__setitem__array-api-strict#207 (comment)
Reacted by Olivier Grisel- state that using disallowed dtype/device combinations raises an error:
__getitem__with the indexer array's device differing from the indexed array device: allowed, prohibited or unspecified?Only unspecified ("implementation-specific") seems like an option here. If PyTorch already does it, can't forbid it (and we don't really have a reason to forbid it anyway). And it definitely shouldn't be mandated that cross-device operations on two arrays need to be supported, that would go against existing design principles.
Same for
__setitem__.Data point:
pytorchallows any combinations, the result's device is the device of the indexed arrayThis is actually surprising by the way, I'd expect a
RuntimeError. I can't verify the claim right now (no GPU at hand) - are you sure you tried with an integer 1-D tensor on a different device, and not with say a 0-D tensor or a list of Python ints?Only unspecified ("implementation-specific") seems like an option here.
+1
An explicit designation as unspecified, implementation-defined would unblock array-api-strict, too.This is actually surprising by the way, I'd expect a RuntimeError
Yes, 1D tensors, and a correction:
s/any combinations/some combinations/g. pytorch apparently special-cases CPU indexers.In [3]: t = torch.arange(9, device='cpu').reshape(3, 3) In [4]: i = torch.arange(3, device='cuda') In [5]: t[i, i] --------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) Cell In[5], line 1 ----> 1 t[i, i] RuntimeError: indices should be either on cpu or on the same device as the indexed tensor (cpu) In [6]: t = torch.arange(9, device='cuda').reshape(3, 3) In [7]: i = torch.arange(3, device='cpu') In [8]: t[i, i] Out[8]: tensor([0, 4, 8], device='cuda:0')- added a commit that references this issue
on Jul 8, 2026
This issue is meant to provide context for changes to dtype support in the next version of the standard.
Overview of data types implemented in various array libraries:
bfloat16in https://docs.jax.dev/en/latest/jax.dtypes.html. See note onfloat64in https://docs.jax.dev/en/latest/default_dtypes.htmlbfloat16support (see https://docs.cupy.dev/en/stable/upgrade.html#minimal-support-for-bfloat16)ml_dtypes: https://github.com/jax-ml/ml_dtypes/blob/main/README.mdSummary of dtype support across array libraries
Legend: ✓ = full support, ○ = partial support, ✗ = no support
Notes:
torch.compile.ml_dtypespackage, not natively.ml_dtypes.bfloat16; some gaps remain especially incupyx.torch.complex32defined, but operator coverage is limited.jax.config.update('jax_enable_x64', True)or theJAX_ENABLE_X64env var; disabled by default, and 64-bit values are silently truncated to 32-bit without it.has_aspect_fp64property is True.has_aspect_fp16property is True.Conclusions
bool,int8/int16/int32,uint8,float32float64is the biggest issue: very important for scientific computing and other fields that require high accuracy, not available at all or CPU-only on several deep learning-focused librariesfloat16andbfloat16don't have universal support yet, but are consistently namedNext steps
Open a PR for discussion which brings documentation on data more in line with reality, reserves the
float16/bfloat16names, and says something about dtype support that is storage-only or partial.