Fix resident loader reuse edge cases

This commit is contained in:
xmarre
2026-04-13 06:16:06 +02:00
parent 9f4b27eb23
commit df9128e9ee
2 changed files with 48 additions and 18 deletions
+1
View File
@@ -554,6 +554,7 @@ def _checkpoint_model_loader_key(
patch_cublaslinear=patch_cublaslinear,
sage_attention=sage_attention,
enable_fp16_accumulation=enable_fp16_accumulation,
extra_state_dict=None,
policy=_effective_policy_name(policy_override),
)
+47 -18
View File
@@ -291,6 +291,40 @@ class ResidencyRegistry:
basename = os.path.basename(source_path) or "anonymous"
return f"{kind}:{basename}:{len(self._entries) + 1}"
def _clear_object_binding(self, obj: Any | None, entry_id: str) -> None:
if obj is None:
return
try:
if self._object_to_entry.get(obj) == entry_id:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
try:
if getattr(obj, "__gpu_resident_loader_entry_id__", None) == entry_id:
delattr(obj, "__gpu_resident_loader_entry_id__")
except (AttributeError, TypeError):
pass
def _tag_object_with_entry(self, obj: Any, entry_id: str) -> None:
try:
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
try:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
return
except (AttributeError, TypeError):
_LOG.debug(
"GPU Resident Loader: could not tag object %r with residency entry id %s",
type(obj),
entry_id,
)
try:
self._object_to_entry[obj] = entry_id
except TypeError:
pass
def bind_object(
self,
obj: Any,
@@ -309,6 +343,8 @@ class ResidencyRegistry:
with self._lock:
old_key: tuple[str, str] | None = None
old_loader_key: tuple[str, str, str] | None = None
previous_obj: Any | None = None
previous_cached_obj: Any | None = None
entry_id = self._loader_key_to_entry.get((kind, source_path, loader_key)) if loader_key is not None else None
if entry_id is None:
entry_id = getattr(obj, "__gpu_resident_loader_entry_id__", None)
@@ -320,6 +356,8 @@ class ResidencyRegistry:
if entry_id is not None and entry_id in self._entries:
entry = self._entries[entry_id]
old_key = (entry.kind, entry.source_path)
previous_obj = entry.object()
previous_cached_obj = entry.cached_object()
if entry.loader_key is not None:
old_loader_key = (entry.kind, entry.source_path, entry.loader_key)
else:
@@ -333,33 +371,22 @@ class ResidencyRegistry:
)
self._entries[entry_id] = entry
self._path_to_entry[(kind, source_path)] = entry_id
try:
setattr(obj, "__gpu_resident_loader_entry_id__", entry_id)
try:
self._object_to_entry.pop(obj, None)
except TypeError:
pass
except (AttributeError, TypeError):
_LOG.debug(
"GPU Resident Loader: could not tag object %r with residency entry id %s",
type(obj),
entry_id,
)
cache_obj = obj if reusable_obj is None else reusable_obj
for stale_obj in (previous_obj, previous_cached_obj):
if stale_obj is None or stale_obj is obj or stale_obj is cache_obj:
continue
self._clear_object_binding(stale_obj, entry_id)
try:
entry.object_ref = weakref.ref(obj)
if getattr(obj, "__gpu_resident_loader_entry_id__", None) is None:
try:
self._object_to_entry[obj] = entry_id
except TypeError:
pass
except TypeError:
entry.object_ref = None
_LOG.debug(
"GPU Resident Loader: object %r is not weak-referenceable; tracking metadata only",
type(obj),
)
cache_obj = obj if reusable_obj is None else reusable_obj
self._tag_object_with_entry(obj, entry_id)
try:
entry.cached_object_ref = weakref.ref(cache_obj)
except TypeError:
@@ -401,6 +428,8 @@ class ResidencyRegistry:
return None
obj = entry.cached_object()
if obj is None:
if entry.cached_object_ref is not None:
return None
obj = entry.object()
if obj is None:
return None