mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-09-12 19:30:38 -07:00
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:
co-authored by
Claude Opus 4.6
parent
a383ff6863
commit
615d604ceb
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user