Add fast_set_attr to modules not inheriting from base.py - #2724
Conversation
Signed-off-by: Varun Thumbe <vthumbe@nvidia.com>
Greptile SummaryThis PR fixes an Key observations:
Confidence Score: 4/5
Important Files Changed
Flowchart%%{init: {'theme': 'neutral'}}%%
flowchart TD
A["prepare_te_modules_for_fsdp(fsdp_root)"] --> B{"_is_te_module(module)?"}
B -->|Yes| C["module.fast_setattr('fsdp_group', process_group)"]
B -->|No| D[Skip]
C --> E{Inherits from\nTransformerEngineBaseModule?}
E -->|Yes - DotProductAttention,\nLinear, LayerNormLinear,\nLayerNormMLP, etc.| F["fast_setattr defined in base.py ✅"]
E -->|No - before this PR| G["AttributeError ❌\nfast_setattr not found"]
E -->|No - after this PR| H["fast_setattr defined in class ✅\nLayerNorm, RMSNorm,\nMultiheadAttention,\nUnfusedDotProductAttention,\nTransformerLayer"]
F --> I["self.__dict__[name] = value\n(bypasses torch.nn.Module.__setattr__)"]
H --> I
Last reviewed commit: d099cfc |
|
/te-ci L1 pytorch |
| def fast_setattr(self, name: str, value: Any) -> None: | ||
| """Fast attribute set for non-parameter fields.""" | ||
| self.__dict__[name] = value |
There was a problem hiding this comment.
Code duplication across 5 new fast_setattr implementations
The same 4-line fast_setattr method is now copy-pasted into LayerNorm, RMSNorm, MultiheadAttention, UnfusedDotProductAttention, and TransformerLayer, in addition to the existing definition in TransformerEngineBaseModule. A lightweight mixin would eliminate this duplication and ensure consistency:
class _FastSetAttrMixin:
def fast_setattr(self, name: str, value: Any) -> None:
"""
Fast version of the Module's set attribute function.
Should be used for regular attributes, but not properties nor parameters/buffers.
"""
self.__dict__[name] = valueThen each class can simply inherit from this mixin. This also reduces the risk of the implementations diverging in the future (e.g., if the base class docstring is updated but the copies are forgotten).
This same duplication exists at transformer_engine/pytorch/module/rmsnorm.py:109-111, transformer_engine/pytorch/attention/multi_head_attention.py:481-483, transformer_engine/pytorch/attention/dot_product_attention/backends.py:296-298, and transformer_engine/pytorch/transformer.py:548-550.
| def fast_setattr(self, name: str, value: Any) -> None: | ||
| """Fast attribute set for non-parameter fields.""" | ||
| self.__dict__[name] = value |
There was a problem hiding this comment.
Incomplete docstring compared to base.py
The docstring here (and in all four other new fast_setattr implementations) is less informative than the one in TransformerEngineBaseModule:
# base.py docstring (full):
Fast version of the Module's set attribute function.
Should be used for regular attributes, but not properties nor parameters/buffers.
# new docstring (incomplete):
Fast attribute set for non-parameter fields.
The key warning — "not properties nor parameters/buffers" — is missing. Because this method bypasses torch.nn.Module.__setattr__ entirely, using it to set a parameter name or a property will silently shadow the parameter/buffer with a plain dict entry rather than updating it correctly, which can cause subtle, hard-to-debug issues. The full constraint should be documented.
| def fast_setattr(self, name: str, value: Any) -> None: | |
| """Fast attribute set for non-parameter fields.""" | |
| self.__dict__[name] = value | |
| def fast_setattr(self, name: str, value: Any) -> None: | |
| """ | |
| Fast version of the Module's set attribute function. | |
| Should be used for regular attributes, but not properties nor parameters/buffers. | |
| """ | |
| self.__dict__[name] = value |
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: