diff --git a/CHANGELOG.md b/CHANGELOG.md index 11c2ee8a5..485c6e9fe 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -75,6 +75,7 @@ Arcade [PyPi Release History](https://pypi.org/project/arcade/#history) page. ### Misc Changes - Ruff's pyupgrade rules are on for `arcade` and `tests`. Type annotations use `list`, `dict`, `X | Y` and `collections.abc` instead of the deprecated `typing` aliases, mostly in `arcade.gl` ([#2676](https://github.com/pythonarcade/arcade/issues/2676)). Nothing changes at runtime. +- Sped up moving sprites in a `SpriteList` with a spatial hash ([#1568](https://github.com/pythonarcade/arcade/issues/1568)). Moving, rotating or resizing a sprite used to remove it from the hash and add it again every time. Now it's only marked as moved, and the next collision check updates each moved sprite once, skipping any still in the same cells. Moving 5,000 hashed sprites with `center_x += 1` and `center_y += 1` went from 36 to 8 ms per frame, or from 39 to 22 ms with a collision check after. Collision checks with nothing moving cost the same. `SpatialHash.contents` and `buckets_for_sprite` are now properties that apply pending moves first. - `draw_lines`, `draw_points`, `draw_line_strip`, `draw_polygon_filled` and `draw_polygon_outline` raise a `ValueError` naming the first point that isn't 2 numbers, such as `point_list[2] is 7, but each point must be 2 numbers, such as (x, y)` ([#2215](https://github.com/pythonarcade/arcade/issues/2215)). Before, a bad point raised a confusing error like `'int' object is not iterable`, and in `draw_lines`, `draw_points` and `draw_line_strip` a point with 1 or 3 numbers silently shifted every number after it. The check costs one length comparison per call, and the points are now converted faster: drawing 1,000 lines or points went from about 137 to 90 µs per call. - The card game tutorial no longer passes `hit_box_algorithm="None"` to `Sprite`, an Arcade 2 argument that `Sprite` silently ignores. Its text no longer says hit box calculation is slow: loading all 52 cards with their default hit boxes takes about 40 ms. - Removed the docs build's workaround for Sphinx not copying changed CSS files (`util/sphinx_static_file_temp_fix.py` and its `.ENABLE_DEVMACHINE_SPHINX_STATIC_FIX` switch). Sphinx fixed it upstream, and the pinned Sphinx 9.1.0 copies changed CSS on incremental builds and with `make.py serve` ([#2266](https://github.com/pythonarcade/arcade/issues/2266)). diff --git a/arcade/sprite_list/spatial_hash.py b/arcade/sprite_list/spatial_hash.py index af8369043..85bb43262 100644 --- a/arcade/sprite_list/spatial_hash.py +++ b/arcade/sprite_list/spatial_hash.py @@ -66,9 +66,14 @@ class SpatialHash(ReadOnlySpatialHash[SpriteType]): """A data structure best for collision checks with non-moving sprites. It subdivides space into a grid of squares, each with sides of length - :py:attr:`cell_size`. Moving a sprite from one place to another is the - same as removing and adding it. Although moving a few can be okay, it - can quickly add up and slow down a game. + :py:attr:`cell_size`. + + Moving a sprite only marks it as moved. The next query, such as a + collision check, puts the moved sprites in their new squares, skipping + any still in the same squares. So a sprite that moves several times in + a frame is updated once, and moving sprites costs nothing until + something checks for collisions. Moving many sprites still adds up and + can slow down a game. Args: cell_size: @@ -91,10 +96,41 @@ def __init__(self, cell_size: int) -> None: width and height. """ # Buckets of sprites per cell - self.contents: dict[IPoint, set[SpriteType]] = {} + self._contents: dict[IPoint, set[SpriteType]] = {} # All the buckets a sprite is in. # This is used to remove a sprite from the spatial hash. - self.buckets_for_sprite: dict[SpriteType, list[set[SpriteType]]] = {} + self._buckets_for_sprite: dict[SpriteType, list[set[SpriteType]]] = {} + # The min and max cells each sprite was added to, to skip moves + # that stay in the same cells + self._cells_for_sprite: dict[SpriteType, tuple[IPoint, IPoint]] = {} + # Sprites that moved since the last query + self._moved: set[SpriteType] = set() + + @property + def contents(self) -> dict[IPoint, set[SpriteType]]: + """The sprites in each cell, keyed by cell coordinates.""" + self._update_moved() + return self._contents + + @property + def buckets_for_sprite(self) -> dict[SpriteType, list[set[SpriteType]]]: + """The cell buckets each sprite is in.""" + self._update_moved() + return self._buckets_for_sprite + + def _update_moved(self) -> None: + """Put sprites that moved since the last query in their new cells.""" + if not self._moved: + return + moved = self._moved + self._moved = set() + cells_for_sprite = self._cells_for_sprite + for sprite in moved: + cells = cells_for_sprite.get(sprite) + # Skip sprites removed since, and moves within the same cells + if cells is not None and self._get_cell_bounds(sprite) != cells: + self.remove(sprite) + self.add(sprite) def hash(self, point: IPoint) -> IPoint: """Convert world coordinates to cell coordinates""" @@ -105,8 +141,10 @@ def hash(self, point: IPoint) -> IPoint: def reset(self): """Clear all the sprites from the spatial hash.""" - self.contents.clear() - self.buckets_for_sprite.clear() + self._contents.clear() + self._buckets_for_sprite.clear() + self._cells_for_sprite.clear() + self._moved.clear() def _get_cell_bounds(self, sprite: BasicSprite) -> tuple[IPoint, IPoint]: """Get the min and max cells covered by a sprite's hit box.""" @@ -123,30 +161,33 @@ def add(self, sprite: SpriteType) -> None: Args: sprite: The sprite to add """ - min_point, max_point = self._get_cell_bounds(sprite) + min_point, max_point = cells = self._get_cell_bounds(sprite) buckets: list[set[SpriteType]] = [] + contents = self._contents # Iterate over the rectangular region adding the sprite to each cell for i in range(min_point[0], max_point[0] + 1): for j in range(min_point[1], max_point[1] + 1): # Add sprite to the bucket - bucket = self.contents.setdefault((i, j), set()) + bucket = contents.setdefault((i, j), set()) bucket.add(sprite) # Collect all the buckets we added to buckets.append(bucket) # Keep track of which buckets the sprite is in - self.buckets_for_sprite[sprite] = buckets + self._buckets_for_sprite[sprite] = buckets + self._cells_for_sprite[sprite] = cells def move(self, sprite: SpriteType) -> None: """ - Shortcut to remove and re-add a sprite. + Mark a sprite as moved. + + It's put in its new cells at the next query, if they changed. Args: sprite: The sprite to move """ - self.remove(sprite) - self.add(sprite) + self._moved.add(sprite) def remove(self, sprite: SpriteType) -> None: """ @@ -156,20 +197,23 @@ def remove(self, sprite: SpriteType) -> None: sprite: The sprite to remove """ # Remove the sprite from all the buckets it is in - for bucket in self.buckets_for_sprite[sprite]: + for bucket in self._buckets_for_sprite[sprite]: bucket.remove(sprite) # Delete the sprite from the bucket tracker - del self.buckets_for_sprite[sprite] + del self._buckets_for_sprite[sprite] + del self._cells_for_sprite[sprite] + self._moved.discard(sprite) # NOTE: The query methods below use contents.get() rather than # setdefault() so that looking at an empty cell doesn't create a bucket # for it. Otherwise the dict grows with every cell ever queried. def get_sprites_near_sprite(self, sprite: BasicSprite) -> set[SpriteType]: + self._update_moved() min_point, max_point = self._get_cell_bounds(sprite) close_by_sprites: set[SpriteType] = set() - contents = self.contents + contents = self._contents # Iterate over the all the covered cells and collect the sprites for i in range(min_point[0], max_point[0] + 1): @@ -181,9 +225,10 @@ def get_sprites_near_sprite(self, sprite: BasicSprite) -> set[SpriteType]: return close_by_sprites def get_sprites_near_point(self, point: Point) -> set[SpriteType]: + self._update_moved() hash_point = self.hash((trunc(point[0]), trunc(point[1]))) # Return a copy of the set. - return set(self.contents.get(hash_point, ())) + return set(self._contents.get(hash_point, ())) def get_sprites_near_rect(self, rect: Rect) -> set[SpriteType]: left, right, bottom, top = rect.lrbt @@ -193,7 +238,8 @@ def get_sprites_near_rect(self, rect: Rect) -> set[SpriteType]: # hash the minimum and maximum points min_point, max_point = self.hash(min_point), self.hash(max_point) close_by_sprites: set[SpriteType] = set() - contents = self.contents + self._update_moved() + contents = self._contents # Iterate over the all the covered cells and collect the sprites for i in range(min_point[0], max_point[0] + 1): @@ -211,4 +257,4 @@ def count(self) -> int: # changing the truthiness of the class instance. # if spatial_hash will be False if it is empty. # For backwards compatibility, we'll keep it as a property. - return len(self.buckets_for_sprite) + return len(self._buckets_for_sprite) diff --git a/benchmarks/spatial_hash/add_remove_vs_move.py b/benchmarks/spatial_hash/add_remove_vs_move.py index 6f7055582..94aef13e4 100644 --- a/benchmarks/spatial_hash/add_remove_vs_move.py +++ b/benchmarks/spatial_hash/add_remove_vs_move.py @@ -17,18 +17,20 @@ def add_remove(): sh = arcade.SpatialHash(CELL_SIZE) for sprite in sprites: - sh.insert_object_for_box(sprite) + sh.add(sprite) for sprite in sprites: - sh.remove_object(sprite) - sh.insert_object_for_box(sprite) + sh.remove(sprite) + sh.add(sprite) def move(): sh = arcade.SpatialHash(CELL_SIZE) for sprite in sprites: - sh.insert_object_for_box(sprite) + sh.add(sprite) for sprite in sprites: sh.move(sprite) + # Moves are applied at the next query + sh.get_sprites_near_point((0, 0)) res_1 = timeit.timeit(add_remove, number=100, globals=globals()) diff --git a/doc/programming_guide/performance_tips.rst b/doc/programming_guide/performance_tips.rst index 029c4d949..e3f014953 100644 --- a/doc/programming_guide/performance_tips.rst +++ b/doc/programming_guide/performance_tips.rst @@ -220,17 +220,23 @@ examples are linked below in :ref:`collision_performance_spatial_hashing_example The Catch """"""""" -Spatial hashing doubles the cost of moving or resizing sprites. +Spatial hashing makes moving, rotating or resizing sprites cost more. -However, this doesn't mean we can't *ever* move or resize a sprite! -Instead, it means we have to be careful about when and how much we -do so. This is because moving and resizing now consists of: +Moving a sprite only marks it as moved. The next collision check puts +every moved sprite back in the right grid squares, which means: -#. Remove it from the internal list of every grid square it is currently in -#. Add it again by re-computing its new location +#. Working out which grid squares the sprite's hit box now covers +#. If they changed, removing it from its old squares and adding it to the + new ones -If we only move a few sprites in the list now and then, it can work out. -When in doubt, test it and see if it works for your specific use case. +So a sprite that moves several times in a frame is only updated once, and +one that stays in the same squares isn't moved at all. Even so, moving +5,000 sprites every frame and then checking for a collision took about +22 ms in a hashed list, compared with 8 ms in an unhashed one. + +This doesn't mean we can't *ever* move a sprite in a hashed list! If we +only move a few sprites now and then, it works out well. When in doubt, +test it and see if it works for your specific use case. .. _collision_performance_spatial_hashing_examples: diff --git a/tests/unit/spritelist/test_spatial_hash_moves.py b/tests/unit/spritelist/test_spatial_hash_moves.py new file mode 100644 index 000000000..4716a021a --- /dev/null +++ b/tests/unit/spritelist/test_spatial_hash_moves.py @@ -0,0 +1,108 @@ +"""Moving a sprite in a spatial hash is applied at the next query.""" + +import arcade +from arcade.sprite_list.spatial_hash import SpatialHash +from arcade.types.rect import LRBT + + +def make_list(*positions): + sprite_list = arcade.SpriteList(use_spatial_hash=True, spatial_hash_cell_size=32) + sprites = [] + for position in positions: + sprite = arcade.SpriteSolidColor(8, 8, color=arcade.color.RED) + sprite.position = position + sprite_list.append(sprite) + sprites.append(sprite) + return sprite_list, sprites + + +def count_adds(spatial_hash: SpatialHash | None): + """Count calls to add() on one spatial hash.""" + assert spatial_hash is not None + calls = [] + original = spatial_hash.add + + def add(sprite): + calls.append(sprite) + original(sprite) + + spatial_hash.add = add + return calls + + +def test_query_finds_moved_sprite(): + sprite_list, (sprite,) = make_list((10, 10)) + spatial_hash = sprite_list.spatial_hash + + sprite.position = 500, 500 + + assert spatial_hash.get_sprites_near_point((10, 10)) == set() + assert spatial_hash.get_sprites_near_point((500, 500)) == {sprite} + assert spatial_hash.get_sprites_near_rect(LRBT(480, 520, 480, 520)) == {sprite} + + sprite.position = 10, 10 + other = arcade.SpriteSolidColor(8, 8) + other.position = 12, 12 + assert spatial_hash.get_sprites_near_sprite(other) == {sprite} + + +def test_collision_check_sees_moved_sprite(): + sprite_list, (sprite,) = make_list((10, 10)) + player = arcade.SpriteSolidColor(16, 16) + player.position = 300, 300 + assert arcade.check_for_collision_with_list(player, sprite_list) == [] + + sprite.position = 302, 298 + assert arcade.check_for_collision_with_list(player, sprite_list) == [sprite] + + +def test_several_moves_update_once(): + sprite_list, (sprite,) = make_list((10, 10)) + adds = count_adds(sprite_list.spatial_hash) + + # Moving the x and y separately, then rotating, used to re-add the + # sprite three times + sprite.center_x = 200 + sprite.center_y = 200 + sprite.angle = 45 + assert adds == [] + + sprite_list.spatial_hash.get_sprites_near_point((200, 200)) + assert adds == [sprite] + + +def test_move_within_same_cells_is_skipped(): + sprite_list, (sprite,) = make_list((10, 10)) + adds = count_adds(sprite_list.spatial_hash) + + sprite.center_x += 1 + assert sprite_list.spatial_hash.get_sprites_near_point((11, 10)) == {sprite} + assert adds == [] + + +def test_remove_after_move(): + sprite_list, (sprite, _other) = make_list((10, 10), (100, 100)) + sprite.position = 500, 500 + sprite_list.remove(sprite) + + spatial_hash = sprite_list.spatial_hash + assert spatial_hash.get_sprites_near_point((500, 500)) == set() + assert spatial_hash.get_sprites_near_point((10, 10)) == set() + assert spatial_hash.count == 1 + + +def test_contents_include_moves(): + sprite_list, (sprite,) = make_list((10, 10)) + sprite.position = 500, 500 + + spatial_hash = sprite_list.spatial_hash + occupied = {cell for cell, bucket in spatial_hash.contents.items() if bucket} + assert occupied == {(15, 15)} + assert spatial_hash.buckets_for_sprite[sprite] == [{sprite}] + + +def test_reset_forgets_moves(): + sprite_list, (sprite,) = make_list((10, 10)) + sprite.position = 500, 500 + sprite_list.spatial_hash.reset() + assert sprite_list.spatial_hash.get_sprites_near_point((500, 500)) == set()