diff --git a/src/outlines/processors/tensor_adapters/numpy.py b/src/outlines/processors/tensor_adapters/numpy.py index 831220444d..5ee17840a3 100644 --- a/src/outlines/processors/tensor_adapters/numpy.py +++ b/src/outlines/processors/tensor_adapters/numpy.py @@ -42,9 +42,7 @@ def boolean_ones_like(self, tensor): return self.numpy.ones_like(tensor, dtype=bool) def apply_mask(self, tensor, mask, value): - result = tensor.copy() - result[mask] = value - return result + return self.numpy.where(mask, value, tensor) def argsort_descending(self, tensor): return self.numpy.argsort(-tensor) diff --git a/tests/processors/test_tensor_adapters.py b/tests/processors/test_tensor_adapters.py index 064afe85d2..b8563ddd87 100644 --- a/tests/processors/test_tensor_adapters.py +++ b/tests/processors/test_tensor_adapters.py @@ -243,6 +243,33 @@ def test_tensor_adapter_apply_mask(framework): assert masked[i, j] == tensor[i, j] +@pytest.mark.parametrize("framework", frameworks) +def test_tensor_adapter_apply_mask_broadcasts_lower_rank_mask(framework): + """A mask with fewer dimensions than the tensor (e.g. the same banned-token + mask reused for every row of a batch) must broadcast, matching torch's + masked_fill semantics, instead of requiring an exact shape match.""" + tensor = create_tensor(framework, (2, 3)) + + if framework == "torch": + mask = torch.tensor([True, False, True]) + elif framework == "numpy": + mask = np.array([True, False, True]) + elif framework == "mlx": + if not HAS_MLX: + pytest.skip("MLX not available") + mask = mx.array([True, False, True]) + + masked = adapters[framework].apply_mask(tensor, mask, float("-inf")) + + assert masked.shape == (2, 3) + for i in range(2): + for j in range(3): + if mask[j]: + assert masked[i, j] == float("-inf") + else: + assert masked[i, j] == tensor[i, j] + + @pytest.mark.parametrize("framework", frameworks) def test_tensor_adapter_argsort_descending(framework): tensor = create_tensor(framework, (2, 3))