spd_learn.functional.recommend_dtype_for_spd#

spd_learn.functional.recommend_dtype_for_spd(condition_number: float, *, prefer_speed: bool = False) → dtype[source]#

Recommend a dtype based on expected matrix condition number.

Parameters:
  • condition_number (float) – The expected condition number of the SPD matrices.

  • prefer_speed (bool, default=False) – If True, prefer faster dtypes when possible.

Returns:

The recommended dtype.

Return type:

torch.dtype

Examples

>>> from spd_learn.functional.numerical import recommend_dtype_for_spd
>>> # Well-conditioned matrices can use float32
>>> print(recommend_dtype_for_spd(1e3))
torch.float32
>>> # Ill-conditioned matrices need float64
>>> print(recommend_dtype_for_spd(1e10))
torch.float64