From e2dbd936e19d27215a27231ac16d99db443df998 Mon Sep 17 00:00:00 2001 From: Siming Chen <867965859@qq.com> Date: Fri, 25 Sep 2026 17:53:42 +0800 Subject: [PATCH] fix: avoid deep-copying the checkpoint in load_pretrained_model --- funasr/train_utils/load_pretrained_model.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/funasr/train_utils/load_pretrained_model.py b/funasr/train_utils/load_pretrained_model.py index 383c22a4b8..042887d8de 100644 --- a/funasr/train_utils/load_pretrained_model.py +++ b/funasr/train_utils/load_pretrained_model.py @@ -8,7 +8,6 @@ import torch.nn import torch.optim import pdb -import copy def load_pretrained_model( @@ -41,7 +40,10 @@ def load_pretrained_model( buffer = BytesIO(oss_bucket.get_object(path).read()) ori_state = torch.load(buffer, map_location=map_location) - src_state = copy.deepcopy(ori_state) + # ori_state is created by torch.load inside this function, is not shared + # with the caller, and is not read again below; src_state is only read. + # Deep-copying every tensor here doubles peak load memory with no effect. + src_state = ori_state src_state = src_state["state_dict"] if "state_dict" in src_state else src_state src_state = src_state["model_state_dict"] if "model_state_dict" in src_state else src_state src_state = src_state["model"] if "model" in src_state else src_state