Add type hints to helpers.py, hotswap.py, constants.py, integrations.py (consolidates #3448, #3452) - #3529
Conversation
|
Hi @BenjaminBossan — pushed one more commit onto this pooled PR, adding type hints to the two functions in |
Continues the pooled type-hint effort in this PR (precedent: huggingface#3144): - gather_params_ctx: Union[nn.Parameter, Iterable[nn.Parameter]] param, Iterator[None] return - dequantize_bnb_weight: Optional[Any] state param, torch.Tensor return - get_layer_device_map / map_cache_to_layer_device_map: Any model param (avoids new pyright errors from nn.Module.__getattr__ ambiguity on dynamically-set hf_device_map/config attributes), transformers.Cache for cache param - init_empty_weights, _init_on_device, _skip_init_on_device: Iterator[None] return - skip_init_on_device: Callable param and return type No functional changes.
BenjaminBossan
left a comment
There was a problem hiding this comment.
Thanks for pooling the PRs to add more type annotations. I have a few comments, otherwise the PR looks good.
| def gather_params_ctx( | ||
| param: Union[nn.Parameter, Iterable[nn.Parameter]], | ||
| modifier_rank: Optional[int] = 0, | ||
| fwd_module: torch.nn.Module = None, |
There was a problem hiding this comment.
Probably most type checkers accept this, but I'd still rather explicitly type fwd_module as being optional.
|
|
||
|
|
||
| def _update_scaling(lora_module, adapter_name, scaling=None): | ||
| def _update_scaling(lora_module: LoraLayer, adapter_name: str, scaling: Optional[float] = None) -> None: |
There was a problem hiding this comment.
Hmm, I think scaling can't ever be None, the default argument here doesn't really make sense. In all PEFT hotswapping tests, it's never None and I also don't see how that could be achieved. So let's remove the default value and type it as float.
| # adapted from: | ||
| # https://github.com/huggingface/transformers/blob/eab6c491d439e83d5e31c660df6f7e36592eb0a2/src/transformers/generation/utils.py#L1617-L1643 | ||
| def get_layer_device_map(model): | ||
| def get_layer_device_map(model: Any) -> Optional[dict[int, Union[int, str]]]: |
There was a problem hiding this comment.
Let's type the output as dict[int, Union[int, str]] | None, which IMO makes more sense for outputs.
| def init_empty_weights(include_buffers: bool | None = None) -> Iterator[None]: | ||
| # adapted from accelerate.big_modeling.py | ||
| with _init_on_device(torch.device("meta"), include_buffers=include_buffers) as f: | ||
| yield f |
There was a problem hiding this comment.
Hmm, yielding f doesn't really make sense here, does it? Let's remove it, then the output type makes more sense.
…tswap.py - gather_params_ctx: fwd_module now Optional[torch.nn.Module] to match its None default - get_layer_device_map: return type uses dict[int, Union[int, str]] | None - init_empty_weights: drop the meaningless yield f (always None), yield bare instead - _update_scaling: scaling is never None in practice, drop Optional and the default
|
Thanks for the review, @BenjaminBossan! Addressed all 4 comments in commit 610410b:
|
BenjaminBossan
left a comment
There was a problem hiding this comment.
Thanks for the updates, LGTM.
What does this PR do?
Consolidates type-hint PRs into a single PR, per @BenjaminBossan's request on #3452:
This PR supersedes and replaces:
rescale_adapter_scaleandcompute_lossinhelpers.pyhotswap.pyBoth #3448 and #3452 will be closed in favor of this PR. No functional behavior is changed in any of the touched files; only type annotations are added to existing signatures (plus one pure variable rename needed to keep a new annotation type-correct, see below).
src/peft/helpers.pyrescale_adapter_scale(model, multiplier)->rescale_adapter_scale(model: nn.Module, multiplier: Union[float, int]) -> Iterator[None]MontecloraTrainerMixin.compute_loss(self, model, inputs, return_outputs=False, **kwargs)-> typed withnn.Module,dict[str, Any],bool, and aUnion[torch.Tensor, tuple[torch.Tensor, Any]]return typesrc/peft/utils/hotswap.py_update_scaling(lora_module, adapter_name, scaling=None)-> typed withLoraLayer,str,Optional[float],-> Nonehotswap_adapter_from_state_dict(...)-> added missing-> Nonereturn annotationhotswap_adapter(model, model_name_or_path, adapter_name, torch_device=None, **kwargs)-> typed withtorch.nn.Module,str,str,Optional[str],-> NoneThis also folds in the fix @githubnemo requested on #3448 (
make style): theIteratorimport now comes fromcollections.abcinstead oftyping, per thedeprecated-importruff rule.src/peft/utils/constants.pyContinuing the same pooled effort (precedent: #3144, merged), the following two functions were missing all parameter and return type hints:
bloom_model_postprocess_past_key_value(past_key_values)->bloom_model_postprocess_past_key_value(past_key_values: tuple[torch.Tensor, ...]) -> tuple[tuple[torch.Tensor, torch.Tensor], ...]starcoder_model_postprocess_past_key_value(past_key_values)->starcoder_model_postprocess_past_key_value(past_key_values: tuple[torch.Tensor, ...]) -> tuple[torch.Tensor, ...]Types were derived from the call site in
peft_model.py(past_key_valuesis produced byTensor.split(...), i.e. atuple[torch.Tensor, ...]) and from tracing each function body.Adding the
tuple[torch.Tensor, ...]parameter annotation tobloom_model_postprocess_past_key_valuesurfaced a real pyright error: the function body reassigned thepast_key_valuesparameter to the result oftorch.cat(past_key_values)(a plainTensor), which is incompatible with the new declared parameter type. Fixed by renaming that local variable toconcatenated_past_key_values— a pure rename with no behavior change, same category of fix as theloftq_utils.py/integrations.pypyright cleanups in #3144.src/peft/utils/integrations.py(new in this update)The remaining functions in this file without full type hints, following the same precedent as #3144 (which added hints to
merge_utils.py/other.pyand did a small pyright fix in this same file):gather_params_ctx(param, modifier_rank=0, fwd_module=None)->param: Union[nn.Parameter, Iterable[nn.Parameter]],-> Iterator[None](it's a@contextmanager; param type derived from the ~15 real call sites acrosstuners/, which pass singlenn.Parameters,.parameters()iterators, and lists)dequantize_bnb_weight(weight, state=None)->state: Optional[Any],-> torch.Tensorget_layer_device_map(model)->model: Any,-> Optional[dict[int, Union[int, str]]]map_cache_to_layer_device_map(model, cache) -> None->model: Any,cache: transformers.Cacheinit_empty_weights,_init_on_device,_skip_init_on_device->-> Iterator[None](all@contextmanager)skip_init_on_device(func)->func: Callable,-> Callablemodelinget_layer_device_map/map_cache_to_layer_device_mapis typedAnyrather thantorch.nn.Module: both functions readmodel.hf_device_mapandmodel.config.num_hidden_layers, attributes thataccelerate/transformersattach dynamically and that aren't part ofnn.Module's stub. Typingmodelastorch.nn.Modulethere resolves those accesses throughnn.Module.__getattr__'sTensor | Modulestub return type and introduces 13 new pyright errors that don't reflect real bugs (same category of false positive#3452hit withLoraLayervstorch.nn.Module);Anyavoids that noise while still documenting the return type precisely.Typing
dequantize_bnb_weight's return astorch.Tensoralso surfaces one pre-existing, real pyright finding indequantize_module_weight(declared-> torch.nn.Parameter, but one code path returns whateverdequantize_bnb_weightreturns, i.e. a plainTensor) — that function already has complete hints from before this PR, so it's left unchanged and out of scope here, but flagging it for visibility.Coordination / approval
make stylefix: Add type hints to rescale_adapter_scale and compute_loss in helpers.py #3448 (review)constants.pyandintegrations.pyadditions follow the same precedent and pooling request rather than opening separate PRs.Testing
ruff checkandruff format --check, run via the repo's pinnedruff~=0.15.12(i.e.make quality's lint/format steps), pass on all four changed files.npx pyright src/peft/utils/constants.py-> 0 errors, 0 warnings after the variable rename described above.npx pyright src/peft/utils/integrations.py-> 27 errors both before and after this change save for the one pre-existing finding indequantize_module_weightdescribed above (26 -> 27); none are newly introduced by theAny-typedmodelparameters.python3 -m py_compileand an isolated module load (importlib, bypassing the package'saccelerateimport chain, which isn't installed in this sandbox) both pass; function signatures verified viainspect.signature.constants.py), so no new unit tests were added; existing tests are unaffected by annotation-only changes.AI disclosure
AI assistance (Claude Code) was used to prepare this PR, including this update: adding type hints to
integrations.pyby reading each function body and its call sites, runningruff/pyrighton the changed file and diffing the error output against the pre-change baseline to confirm no new false positives, and documenting the one genuine pre-existing pyright finding the new annotations surfaced. All changes were reviewed by the submitter (@RudrenduPaul) before pushing.