Repository navigation
Fix ty suppressions from #405 with real type fixes (#408) - #769
lucianlavric wants to merge 5 commits into
Conversation
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
left a comment
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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), ( |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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] |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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.
| # 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 |
There was a problem hiding this comment.
I have a preference against declaration only statements like this, can we avoid these?
There was a problem hiding this comment.
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.
|
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.
Summary
#405 introduced
ty: ignoresuppressions (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_inputsdeclareddecoder: LinkPredictionDecoderbut assigns the bounddecodemethod; it is now typed as the callable it actually is.Approach
LightGCN._layer_weights, the early-stop test's_DummyModel.foo) so ty resolves them asTensorrather than theTensor | Moduleunion fromnn.Module.__getattr__.GnnModelprotocolruntime_checkableand narrow through it inInferencerV1; keep the dynamicgraph_backendattach in the GraphSAGE spec explicit viasetattr.infer_task_inputs(LinkPredictionGNN), andModuleDictvalues inNodeAnchorBasedLinkPredictionTasks._get_all_tasks.FeatureInteraction.reset_parameters: bind viagetattrsocallable()narrows the optional method.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-onlyos/psutilAPIs).make check_format_py: clean.make unit_test_pyis not runnable on macOS (graphlearn_torchships in the Docker/dev images only; the darwin env lackstensorflow_metadata/torchrec), unchanged from main — relying on PR CI for the test gate.+=behavior unchanged with the declared annotation;runtime_checkableGnnModelisinstance correctly passes aftersetattrand fails on modules without the attribute; the explicitly importedspawnis callable.