fix(numpy-compat): raise on unknown dtype + add fp16/complex

Follow-up to #361. The original fallback silently mapped unknown numpy
dtypes to torch.float32, which would reinterpret the memcpy'd bytes in
the wrong dtype and corrupt data (e.g. fp16 tensors from some TTS
engines) rather than erroring loudly.

- Hoist dtype_map out of the inner function so it's built once
- Add float16, complex64, complex128 mappings
- Raise TypeError on unknown dtype instead of silent float32 fallback

Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
This commit is contained in:
James Pine
2026-04-16 01:47:22 -07:00
co-authored by Claude Opus 4.6
parent a383ff6863
commit 615d604ceb
+30 -15
View File
@@ -48,25 +48,40 @@ def _patch_torch_from_numpy():
_orig = torch.from_numpy
def _safe_from_numpy(arr, _orig=_orig, _c=ctypes, _np=np, _t=torch):
# Explicit numpy → torch dtype map. Silent fallback to float32 on
# unknown dtypes would reinterpret the memcpy'd bytes as fp32 and
# silently corrupt data (e.g. fp16 tensors from some TTS engines),
# so we raise instead.
dtype_map = {
"float16": _t.float16,
"float32": _t.float32,
"float64": _t.float64,
"int8": _t.int8,
"int16": _t.int16,
"int32": _t.int32,
"int64": _t.int64,
"uint8": _t.uint8,
"bool": _t.bool,
"complex64": _t.complex64,
"complex128": _t.complex128,
}
def _safe_from_numpy(
arr, _orig=_orig, _c=ctypes, _np=np, _t=torch, _map=dtype_map
):
try:
return _orig(arr)
except RuntimeError:
a = _np.ascontiguousarray(arr)
dtype_map = {
"float32": _t.float32,
"float64": _t.float64,
"int32": _t.int32,
"int64": _t.int64,
"int16": _t.int16,
"int8": _t.int8,
"uint8": _t.uint8,
"bool": _t.bool,
}
out = _t.empty(
list(a.shape),
dtype=dtype_map.get(str(a.dtype), _t.float32),
)
key = str(a.dtype)
if key not in _map:
raise TypeError(
f"pyi_rth_numpy_compat: unsupported numpy dtype "
f"{key!r} in torch.from_numpy fallback; add an "
f"explicit mapping rather than silently copying "
f"bytes into the wrong dtype."
)
out = _t.empty(list(a.shape), dtype=_map[key])
_c.memmove(out.data_ptr(), a.ctypes.data, a.nbytes)
return out