diff --git a/src/deepwave/location_interpolation.py b/src/deepwave/location_interpolation.py index dc30377..bf9aef8 100644 --- a/src/deepwave/location_interpolation.py +++ b/src/deepwave/location_interpolation.py @@ -68,8 +68,15 @@ def _get_hicks_for_one_location_dim( - location + int(location) ) - locations = (location + x).long() - + locations = ( + torch.arange( + -halfwidth + 1, + halfwidth + 1, + device=beta.device, + ) + + int(location) + ) + if key in hicks_weight_cache: weights = hicks_weight_cache[key].clone() else: diff --git a/tests/test_location_interpolation.py b/tests/test_location_interpolation.py index af3fb2e..943173a 100644 --- a/tests/test_location_interpolation.py +++ b/tests/test_location_interpolation.py @@ -389,3 +389,30 @@ def test_get_hicks_for_one_location_dim_invalid_n_grid_points() -> None: [-0.5, 19.5], -1, ) + + +@pytest.mark.parametrize( + "location", + [3.22, 7.35, 15.2, 31.4, 63.16, 127.23], +) +def test_hicks_locations_not_truncated_by_float32(location, halfwidth=2): + """Grid indices must not be truncated by float32 round-off. + + The old formula ``(location + x).long()`` can yield 1.999... for the + leftmost cell, which ``.long()`` truncates (e.g. 1 instead of 2 at + location 3.01). Indices must be integer arithmetic around ``int(location)``. + """ + betas = [0.0, 1.84, 3.04, 4.14, 5.26, 6.40, 7.51, 8.56, 9.56, 10.64] + beta = torch.tensor(betas[halfwidth - 1], dtype=torch.float32) + locs, _ = _get_hicks_for_one_location_dim( + {}, + location, + halfwidth, + beta, + [False, False], + [-0.5, 199.5], + 200, + True, + ) + expected = torch.arange(-halfwidth + 1, halfwidth + 1) + int(location) + assert torch.equal(locs.cpu(), expected)