Skip to content

Commit b8c5986

Browse files
committed
Use weak references to resolve circular references and add test
1 parent 670bbee commit b8c5986

2 files changed

Lines changed: 68 additions & 2 deletions

File tree

ignite/engine/engine.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -328,7 +328,7 @@ def execute_something():
328328

329329
try:
330330
_check_signature(handler, "handler", self, *(event_args + args), **kwargs)
331-
self._event_handlers[event_name].append((handler, (self,) + args, kwargs))
331+
self._event_handlers[event_name].append((handler, (weakref.ref(self),) + args, kwargs))
332332
except ValueError:
333333
_check_signature(handler, "handler", *(event_args + args), **kwargs)
334334
self._event_handlers[event_name].append((handler, args, kwargs))
@@ -432,7 +432,15 @@ def _fire_event(self, event_name: Any, *event_args: Any, **event_kwargs: Any) ->
432432
self.last_event_name = event_name
433433
for func, args, kwargs in self._event_handlers[event_name]:
434434
kwargs.update(event_kwargs)
435-
first, others = ((args[0],), args[1:]) if (args and args[0] == self) else ((), args)
435+
if args and isinstance(args[0], weakref.ref):
436+
resolved_engine = args[0]()
437+
if resolved_engine is None:
438+
raise RuntimeError("Engine reference not resolved. Cannot execute event handler.")
439+
first, others = ((resolved_engine,), args[1:])
440+
else:
441+
# metrics do not provide engine when registered
442+
first, others = (tuple(), args) # type: ignore[assignment]
443+
436444
func(*first, *(event_args + others), **kwargs)
437445

438446
def fire_event(self, event_name: Any) -> None:
Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,58 @@
1+
import weakref
2+
3+
import pytest
4+
5+
from ignite.engine import Engine, Events
6+
7+
8+
class TestEngineMemoryLeak:
9+
"""See: https://github.com/pytorch/ignite/issues/3438"""
10+
11+
ENGINE_WEAK_REFS = set()
12+
13+
def do_train(self, cls, with_handler) -> None:
14+
engine = cls(lambda e, b: None)
15+
16+
if with_handler:
17+
18+
@engine.on(Events.EPOCH_STARTED)
19+
def handler(engine) -> None:
20+
pass
21+
22+
engine.run(range(5), max_epochs=5)
23+
self.ENGINE_WEAK_REFS.add(weakref.ref(engine))
24+
25+
@pytest.mark.parametrize(
26+
"with_handler",
27+
[
28+
(True,),
29+
(False,),
30+
],
31+
)
32+
def test_memory_leak(self, with_handler):
33+
num_iters = 5
34+
counter = 0
35+
36+
class EngineForTests(Engine):
37+
38+
def __del__(self):
39+
nonlocal counter
40+
counter += 1
41+
42+
for i in range(num_iters):
43+
self.do_train(EngineForTests, with_handler)
44+
for weak_engine_ref in self.ENGINE_WEAK_REFS:
45+
engine = weak_engine_ref()
46+
assert engine is None
47+
48+
print(counter)
49+
if with_handler:
50+
assert counter == i + 1
51+
else:
52+
assert counter == 0
53+
54+
print(counter)
55+
if with_handler:
56+
assert counter == i + 1
57+
else:
58+
assert counter == 0

0 commit comments

Comments
 (0)