Skip to content

fix: avoid deep-copying the checkpoint in load_pretrained_model - #3729

Merged
LauraGPT merged 1 commit into
modelscope:mainfrom
yydhYYDH:perf/avoid-deepcopy-in-load-pretrained-model
Sep 26, 2026
Merged

LauraGPT merged 1 commit into
modelscope:mainfrom
yydhYYDH:perf/avoid-deepcopy-in-load-pretrained-model

Conversation

@yydhYYDH

Copy link
Copy Markdown
Contributor

Related to #3728

Summary

load_pretrained_model() deep-copies the whole checkpoint state dict on every load. That copy is redundant:

  • ori_state is created by torch.load inside the function, so it has a single owner and is not shared with the caller.
  • ori_state is not referenced again after the deepcopy line.
  • src_state is only read from that point on. The loop inspects src_state[k_src].shape and rebinds references in dst_state, and load_state_dict is what copies data into the model. No tensor is mutated in place.

So the deep copy changes nothing while holding a second full copy of the checkpoint in memory for the whole load. Peak host memory during load drops by roughly the size of the checkpoint: for iic/speech_paraformer-large_asr_nat-zh-cn-16k-common-vocab8404-online (220M params, 956 tensors, 840MB checkpoint) peak RSS goes from 3461 MB to 2624 MB, a 837MB saving that matches the checkpoint size. That matters on memory-constrained hosts and under container memory limits, where the current peak can fail a load that would otherwise fit.

The now-unused import copy is removed as well.

Type of change

  • Bug fix
  • Documentation
  • Example or demo
  • Runtime or deployment
  • Benchmark or evaluation
  • Model/training change

Validation

  • python -m compileall funasr examples tests
  • Docs or links checked
  • Runtime/deployment command tested

Same machine, same checkpoint, alternating runs, via AutoModel(..., device="cpu", disable_update=True):

funasr 1.4.16     : peak RSS 3474 / 3462 MB   <All keys matched successfully>
with this change : peak RSS 2635 / 2623 MB   <All keys matched successfully>

The ~840MB reduction is stable across runs and tracks the checkpoint size, which is the copy being dropped. All keys matched successfully on both sides is the correctness signal: the same parameters are loaded either way.

Wall-clock is not quantified here: the host runs other builds and tests concurrently, so the load time varies with unrelated load. Directionally the copy is also faster to skip, but treat the memory figure as the reproducible result.

Runtime check on this branch after the change, streaming ASR at 600ms chunks on an RTX 4060:

loaded 44.3s | VRAM 851 MB
audio 5.55s -> 10 chunks x 600ms
RTF=0.2588 | TEXT='欢迎大家来体验达摩院推出的语音识别模型'

User impact

Anyone loading a FunASR model — particularly on hosts with a tight memory budget, in containers with memory limits, or on the edge deployments the toolkit targets. The saving scales with checkpoint size, so it is most visible on the large ASR and LLM-ASR checkpoints.

Notes for reviewers

  • No API, signature, or checkpoint-format change. The diff is one file, +4/-2.
  • The reasoning is in the commit message; the short version is that ori_state has exactly one owner and is dead after the copy, and src_state is read-only downstream.
  • If a future caller genuinely needs ori_state to survive the call, that would be a different contract. It is not the case today, and the variable is local to this function.
  • Happy to add a regression test that pins load-time peak memory. The measurement above is resource.getrusage-based rather than a unit test, so it is reported as a benchmark rather than a CI assertion.

@LauraGPT LauraGPT 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.

Verified exact head e2dbd93 against current main 41778c4. The loader only reads the selected source state before load_state_dict; focused repository tests pass 13/13, independent direct/wrapped/DDP/mismatch contract probes pass 6/6 on both base and head, and a separate-process 128 MiB checkpoint probe reduces peak RSS by about 143 MiB while preserving loaded values. The merge tree is conflict-free; compileall and diff-check pass.

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