Skip to content

Fix ty suppressions from #405 with real type fixes (#408) - #769

Open
lucianlavric wants to merge 5 commits into
Snapchat:mainfrom
lucianlavric:luka/408-ty-union-inference-fixes
Open

lucianlavric wants to merge 5 commits into
Snapchat:mainfrom
lucianlavric:luka/408-ty-union-inference-fixes

Conversation

@lucianlavric

Copy link
Copy Markdown

Summary

#405 introduced ty: ignore suppressions (tracked in #408) where the type fixes were not trivial. This removes every #408-linked suppression by fixing the underlying types instead: declared buffer types instead of __getattr__ unions, fail-fast narrowing instead of ignores, and one real mismatch — infer_task_inputs declared decoder: LinkPredictionDecoder but assigns the bound decode method; it is now typed as the callable it actually is.

Approach

  • Declare attribute types for registered buffers (LightGCN._layer_weights, the early-stop test's _DummyModel.foo) so ty resolves them as Tensor rather than the Tensor | Module union from nn.Module.__getattr__.
  • Make the GnnModel protocol runtime_checkable and narrow through it in InferencerV1; keep the dynamic graph_backend attach in the GraphSAGE spec explicit via setattr.
  • Narrow with fail-fast asserts where the runtime type is guaranteed by construction: the model in infer_task_inputs (LinkPredictionGNN), and ModuleDict values in NodeAnchorBasedLinkPredictionTasks._get_all_tasks.
  • FeatureInteraction.reset_parameters: bind via getattr so callable() narrows the optional method.
  • KDD tutorial: from torch.multiprocessing.spawn import spawn — the dotted call resolves to the same-named submodule, which type checkers reject as non-callable.

User-facing changes

None intended. Invalid model types now fail fast with descriptive assert messages instead of failing deeper in the call.

Verification

  • make type_check: no new diagnostics versus main (the same 48 pre-existing macOS-environment diagnostics appear on both — uninstallable-on-macOS imports and two Linux-only os/psutil APIs).
  • make check_format_py: clean.
  • Full make unit_test_py is not runnable on macOS (graphlearn_torch ships in the Docker/dev images only; the darwin env lacks tensorflow_metadata/torchrec), unchanged from main — relying on PR CI for the test gate.
  • Runtime patterns validated standalone with pure torch: buffer registration and += behavior unchanged with the declared annotation; runtime_checkable GnnModel isinstance correctly passes after setattr and fails on modules without the attribute; the explicitly imported spawn is callable.

Resolves the ty suppressions tracked in Snapchat#408 by fixing the underlying
types instead of ignoring them:

- Declare buffer types (models.py LightGCN._layer_weights, early stop
  test _DummyModel.foo) so ty resolves them as Tensor rather than the
  Tensor | Module union from nn.Module.__getattr__.
- infer_task_inputs: narrow the model to LinkPredictionGNN with a
  fail-fast assert, and type decoder as the bound decode callable (its
  previous LinkPredictionDecoder declaration did not match the assigned
  bound method).
- Make the GnnModel protocol runtime_checkable and narrow through it in
  InferencerV1 instead of ignoring; keep the dynamic graph_backend
  attach in GraphSAGE spec explicit via setattr.
- NodeAnchorBasedLinkPredictionTasks._get_all_tasks: assert the
  ModuleDict value is a NodeAnchorBasedLinkPredictionBaseTask.
- FeatureInteraction.reset_parameters: bind via getattr so callable()
  narrows the optional method.
- KDD tutorial: import torch.multiprocessing's spawn function explicitly
  (the dotted path resolves to the same-named submodule).

make type_check reports no new diagnostics (remaining are pre-existing
macOS environment gaps also present on main).

@kmontemayor2-sc kmontemayor2-sc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for removing the ty suppressions, and sorry for the late review 😅 . I found one blocking runtime regression in the current implementation: the new LinkPredictionGNN import creates a circular dependency that prevents the node-anchor link-prediction modeling spec from importing.


from gigl.src.common.models.layers.decoder import LinkPredictionDecoder
from gigl.src.common.models.layers.loss import ModelResultType
from gigl.src.common.models.pyg.link_prediction import LinkPredictionGNN

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This import introduces a circular dependency:

utils/infer.py
  → models/pyg/link_prediction.py
  → models/layers/task.py
  → utils/infer.py

The same import succeeds at the PR’s base commit. This blocks NodeAnchorBasedLinkPredictionModelingTaskSpec from loading.
Please avoid importing the concrete model here at runtime. A small structural Protocol describing the required decode and tasks.result_types members would both avoid the cycle and preserve the existing model contract. A TYPE_CHECKING-only import plus a cast would also avoid the cycle, although it would not provide runtime validation.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Confirmed the cycle (infer → link_prediction → task → infer) and removed the import. infer_task_inputs now narrows against a module-private runtime_checkable Protocol declaring decode and tasks.result_types, defined in infer.py, so no concrete model is imported. Verified NodeAnchorBasedLinkPredictionModelingTaskSpec imports again.

if isinstance(model, torch.nn.parallel.DistributedDataParallel)
else model
)
assert isinstance(base_model, LinkPredictionGNN), (

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This changes the function from accepting any compatible nn.Module to requiring one specific legacy LinkPredictionGNN class.
Previously, a custom model or wrapper worked as long as it exposed decode and tasks.result_types, which matches the function’s generic nn.Module | DDP signature. The new check rejects compatible custom implementations, and the required concrete class is itself marked deprecated.
Could we type this against a structural protocol instead? If runtime validation is required, please raise an explicit TypeError rather than using assert, since the assertion disappears under python -O.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Switched to the structural Protocol, so any module exposing decode and tasks.result_types is accepted again, and replaced the assert with raise TypeError. Made the same assert → TypeError change to the two other guards this PR introduced (InferencerV1 and _get_all_tasks) for consistency.

batch_result_types: Set[ModelResultType]
decoder: LinkPredictionDecoder
# The model's bound decode method: (query_embeddings, candidate_embeddings) -> scores.
decoder: Callable[[torch.Tensor, torch.Tensor], torch.Tensor]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This declaration is safe because decoder is assigned unconditionally below, but I would initialize it where it is declared rather than leave a declaration-only local:

decoder: Callable[
    [torch.Tensor, torch.Tensor], torch.Tensor
] = base_model.decode

Alternatively, if the narrowed model type exposes the correct signature, decoder = base_model.decode should be sufficient.
I would not add a None or no-op default here: the decoder is required and model-specific, so such a default would weaken the invariant and could hide an invalid model.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dropped the declaration; decoder and batch_result_types are inferred from the Protocol now. One wrinkle: after the isinstance narrowing ty types the value as Module & _LinkPredictionModel and resolves attributes through Module.getattr as well, giving set | Tensor | Module, so a cast to the Protocol after the runtime check is still needed. Same reason InferencerV1 casts to GnnModel. Comment in the code explains it.

Comment thread gigl/nn/models.py Outdated
# Register layer weights as a buffer so it moves with the model to different devices
# Declared type lets ty resolve the buffer as a Tensor instead of the
# Tensor | Module union produced by nn.Module.__getattr__.
self._layer_weights: torch.Tensor

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I have a preference against declaration only statements like this, can we avoid these?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Replaced the annotation + register_buffer pair with self._layer_weights = nn.Buffer(...) (torch ≥ 2.4, repo pins 2.8). Assignment registers the buffer and ty infers Tensor from the value. Same change to _DummyModel.foo in early_stop_test.py. Checked state_dict keys, persistence, requires_grad, and in-place += are unchanged; early-stop tests pass.

…l and nn.Buffer

- infer.py: replace the LinkPredictionGNN isinstance narrowing with a
  module-private runtime_checkable Protocol (_LinkPredictionModel with
  decode and tasks.result_types). Importing the concrete model created
  the cycle infer -> link_prediction -> task -> infer, which stopped
  NodeAnchorBasedLinkPredictionModelingTaskSpec from importing. The
  Protocol also restores the original contract: any module exposing
  those members is accepted, not only the deprecated LinkPredictionGNN.
- infer.py: raise TypeError instead of assert so the guard survives
  python -O; drop the declaration-only decoder and batch_result_types
  locals now that both are inferred through the Protocol. A cast to the
  Protocol after the runtime check is still required because ty resolves
  attributes on the narrowed Module & Protocol type through
  Module.__getattr__ as well.
- gnn_inferencer.py, task.py: same assert -> raise TypeError change for
  the other two guards this PR introduced.
- models.py, early_stop_test.py: register buffers by assigning nn.Buffer
  instead of a declaration-only annotation plus register_buffer. State
  dict keys, persistence, requires_grad, and in-place updates are
  unchanged.
@lucianlavric

Copy link
Copy Markdown
Author

Addressed all four. The blocking cycle is fixed and verified by importing the node-anchor spec. make type_check matches main's baseline and make check_format_py is clean. Ready for another look.

Resolve conflicts by taking main's versions of graphsage_template_modeling_spec.py,
types/model.py, and gnn_inferencer.py: Snapchat#770 removed GraphBackend, GnnModel, and the
graph-builder factory plumbing, so this branch's changes to those files no longer
apply.
@lucianlavric
lucianlavric marked this pull request as ready for review September 28, 2026 03:39

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants