Skip to content

Commit 8953fe2

Browse files
tillahoffmannclaude
andcommitted
Revert dual-backend Constraint.__call__ returns to ArrayLike
The np/jnp constraint methods (_CorrCholesky, _LowerCholesky, _L1Ball, _PositiveDefinite, _PositiveSemiDefinite, _IndependentConstraint) can return a numpy array on the numpy backend, so ty infers an `np.ndarray | Array` union that the narrowed `-> Array` annotation rejects. A blanket `# type: ignore` was unstable across numpy stub versions: needed under numpy 2.5.0 (py3.14) but flagged unused under numpy 2.4.6 (py3.11), failing lint on one Python version or the other. Revert these overrides plus the base Constraint.__call__/check and the two _validate_sample passthroughs to ArrayLike, which typechecks under both numpy versions. The ~20 pure-jnp constraint overrides keep their narrowed -> Array (a valid covariant override). Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
1 parent 3cc3120 commit 8953fe2

2 files changed

Lines changed: 12 additions & 13 deletions

File tree

‎numpyro/distributions/constraints.py‎

Lines changed: 10 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -97,13 +97,13 @@ def __init_subclass__(cls, **kwargs):
9797
def tree_flatten(self):
9898
raise NotImplementedError
9999

100-
def __call__(self, x: NumLikeT) -> Array:
100+
def __call__(self, x: NumLikeT) -> ArrayLike:
101101
raise NotImplementedError
102102

103103
def __repr__(self) -> str:
104104
return self.__class__.__name__[1:] + "()"
105105

106-
def check(self, value: NumLikeT) -> Array:
106+
def check(self, value: NumLikeT) -> ArrayLike:
107107
"""
108108
Returns a byte tensor of `sample_shape + batch_shape` indicating
109109
whether each event in value satisfies this constraint.
@@ -171,7 +171,7 @@ def feasible_like(self, prototype: NumLike) -> NumLike:
171171
class _CorrCholesky(_SingletonConstraint[NonScalarArray]):
172172
event_dim = 2
173173

174-
def __call__(self, x: NonScalarArray) -> Array:
174+
def __call__(self, x: NonScalarArray) -> ArrayLike:
175175
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
176176
tril = xp.tril(x)
177177
lower_triangular = xp.all(xp.reshape(tril == x, x.shape[:-2] + (-1,)), axis=-1)
@@ -387,7 +387,7 @@ def is_discrete(self) -> bool:
387387
def event_dim(self) -> int:
388388
return self.base_constraint.event_dim + self.reinterpreted_batch_ndims
389389

390-
def __call__(self, x: NumLikeT) -> Array:
390+
def __call__(self, x: NumLikeT) -> ArrayLike:
391391
result = self.base_constraint(x)
392392
if self.reinterpreted_batch_ndims == 0:
393393
return result
@@ -641,7 +641,7 @@ def __repr__(self) -> str:
641641
class _LowerCholesky(_SingletonConstraint[NonScalarArray]):
642642
event_dim = 2
643643

644-
def __call__(self, x: NonScalarArray) -> Array:
644+
def __call__(self, x: NonScalarArray) -> ArrayLike:
645645
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
646646
tril = xp.tril(x)
647647
lower_triangular = xp.all(xp.reshape(tril == x, x.shape[:-2] + (-1,)), axis=-1)
@@ -687,7 +687,7 @@ class _L1Ball(_SingletonConstraint[NumLike]):
687687
event_dim = 1
688688
reltol = 10.0 # Relative to finfo.eps.
689689

690-
def __call__(self, x: NumLike) -> Array:
690+
def __call__(self, x: NumLike) -> ArrayLike:
691691
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
692692
dtype = x.dtype if isinstance(x, xp.ndarray) else type(x)
693693
eps = jnp.finfo(dtype).eps
@@ -710,7 +710,7 @@ def feasible_like(self, prototype: NonScalarArray) -> NonScalarArray:
710710
class _PositiveDefinite(_SingletonConstraint[NonScalarArray]):
711711
event_dim = 2
712712

713-
def __call__(self, x: NonScalarArray) -> Array:
713+
def __call__(self, x: NonScalarArray) -> ArrayLike:
714714
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
715715
# check for symmetric
716716
symmetric = xp.all(xp.isclose(x, xp.swapaxes(x, -2, -1)), axis=(-2, -1))
@@ -738,7 +738,7 @@ def feasible_like(self, prototype: NonScalarArray) -> NonScalarArray:
738738
class _PositiveSemiDefinite(_SingletonConstraint[NonScalarArray]):
739739
event_dim = 2
740740

741-
def __call__(self, x: NonScalarArray) -> Array:
741+
def __call__(self, x: NonScalarArray) -> ArrayLike:
742742
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
743743
# check for symmetric
744744
symmetric = xp.all(xp.isclose(x, xp.swapaxes(x, -2, -1)), axis=(-2, -1))
@@ -853,11 +853,10 @@ def event_dim(self) -> int:
853853
def __call__(self, x: NonScalarArray) -> Array:
854854
xp = np if isinstance(x, (np.ndarray, np.generic)) else jnp
855855
tol = xp.finfo(x.dtype).eps * x.shape[-1] * 10
856-
zerosum_true = True
856+
zerosum_true = jnp.asarray(True)
857857
for dim in range(-self.event_dim, 0):
858858
zerosum_true = zerosum_true & xp.allclose(x.sum(dim), 0, atol=tol)
859-
# FIXME: shape must match batch shape of `x`, not be a literal boolean.
860-
return zerosum_true # type: ignore
859+
return zerosum_true
861860

862861
def eq(self, other: object, static: bool = False) -> bool:
863862
if not isinstance(other, _ZeroSum):

‎numpyro/distributions/distribution.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -426,7 +426,7 @@ def variance(self) -> Array:
426426
def mode(self) -> Array:
427427
raise NotImplementedError
428428

429-
def _validate_sample(self, value: ArrayLike) -> Array:
429+
def _validate_sample(self, value: ArrayLike) -> ArrayLike:
430430
assert self.support is not None
431431
mask = self.support(value)
432432
if not_jax_tracer(mask):
@@ -909,7 +909,7 @@ def log_prob(self, value: ArrayLike) -> Array:
909909
batch_shape = lax.broadcast_shapes(batch_shape, self.batch_shape)
910910
return jnp.zeros(batch_shape)
911911

912-
def _validate_sample(self, value: ArrayLike) -> Array:
912+
def _validate_sample(self, value: ArrayLike) -> ArrayLike:
913913
mask = super(ImproperUniform, self)._validate_sample(value)
914914
batch_dim = jnp.ndim(value) - len(self.event_shape)
915915
if batch_dim < jnp.ndim(mask):

0 commit comments

Comments
 (0)