From f06db796abc8fe00eeb11fec3251a6b544f5412f Mon Sep 17 00:00:00 2001 From: nightcityblade Date: Mon, 22 Jun 2026 11:32:57 +0800 Subject: [PATCH] fix: prevent once handlers from refiring on epoch termination Fixes pytorch/ignite#3625 --- ignite/engine/engine.py | 3 +++ ignite/engine/events.py | 5 ++++- tests/ignite/engine/test_engine.py | 1 + 3 files changed, 8 insertions(+), 1 deletion(-) diff --git a/ignite/engine/engine.py b/ignite/engine/engine.py index 98bea1563c1b..ba8800a8bc16 100644 --- a/ignite/engine/engine.py +++ b/ignite/engine/engine.py @@ -325,6 +325,9 @@ def execute_something(): return RemovableEventHandle(event_name, handler, self) if isinstance(event_name, CallableEventWithFilter) and event_name.filter is not None: event_filter = event_name.filter + once = getattr(event_filter, "_once", None) + if once is not None: + event_filter = Events.once_event_filter(list(once)) handler = self._handler_wrapper(handler, event_name, event_filter) self._assert_allowed_event(event_name) diff --git a/ignite/engine/events.py b/ignite/engine/events.py index 0a370e9d8815..ffd2354420af 100644 --- a/ignite/engine/events.py +++ b/ignite/engine/events.py @@ -140,12 +140,15 @@ def wrapper(engine: "Engine", event: int) -> bool: @staticmethod def once_event_filter(once: list) -> Callable: """A wrapper for once event filter.""" + remaining = set(once) def wrapper(engine: "Engine", event: int) -> bool: - if event in once: + if event in remaining: + remaining.remove(event) return True return False + wrapper._once = tuple(once) # type: ignore[attr-defined] return wrapper @staticmethod diff --git a/tests/ignite/engine/test_engine.py b/tests/ignite/engine/test_engine.py index 508d29570939..bcb0cdc0a92c 100644 --- a/tests/ignite/engine/test_engine.py +++ b/tests/ignite/engine/test_engine.py @@ -395,6 +395,7 @@ def check_previous_events2(): epoch_completed_events = [e for e in engine.called_events if e[2] == Events.EPOCH_COMPLETED.name] assert len(epoch_completed_events) == max_epochs - skip_epoch_completed + assert call_count == 1 @pytest.mark.parametrize("data", [None, "mock_data_loader"]) def test_iteration_events_are_fired(self, data):