optm(checkpoint-engine): disable multi-thread loading when update weights (#12374)

Signed-off-by: Yang Kaiyong <yangkaiyong.yky@antgroup.com>
This commit is contained in:
Yang Kaiyong
2025-11-07 21:30:04 +08:00
committed by GitHub
parent c67fce160e
commit 61bfd9fa9b
3 changed files with 110 additions and 30 deletions
+30 -15
View File
@@ -80,6 +80,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
VocabParallelEmbedding,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.utils import maybe_executor_submit, should_async_load
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
from sglang.srt.server_args import get_global_server_args
@@ -854,6 +855,7 @@ class LongcatFlashForCausalLM(nn.Module):
params_dict = dict(self.named_parameters())
weight_names = []
for name, loaded_weight in weights:
use_async_loading = should_async_load(loaded_weight)
if "mtp" in name:
continue
weight_names.append(name)
@@ -877,8 +879,12 @@ class LongcatFlashForCausalLM(nn.Module):
continue
param = params_dict[name]
weight_loader = param.weight_loader
futures.append(
executor.submit(weight_loader, param, loaded_weight, shard_id)
maybe_executor_submit(
executor=executor,
futures=futures,
use_async=use_async_loading,
func=weight_loader,
func_args=(param, loaded_weight, shard_id),
)
break
else:
@@ -889,15 +895,16 @@ class LongcatFlashForCausalLM(nn.Module):
name = name.replace(weight_name, param_name)
param = params_dict[name]
weight_loader = param.weight_loader
futures.append(
executor.submit(
weight_loader,
param,
loaded_weight,
name,
shard_id=shard_id,
expert_id=expert_id,
)
maybe_executor_submit(
executor=executor,
futures=futures,
use_async=use_async_loading,
func=weight_loader,
func_args=(param, loaded_weight, name),
func_kwargs={
"shard_id": shard_id,
"expert_id": expert_id,
},
)
break
else:
@@ -951,8 +958,12 @@ class LongcatFlashForCausalLM(nn.Module):
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
futures.append(
executor.submit(weight_loader, param, fused_weight)
maybe_executor_submit(
executor=executor,
futures=futures,
use_async=use_async_loading,
func=weight_loader,
func_args=(param, fused_weight),
)
cached_a_proj.pop(q_a_proj_name)
cached_a_proj.pop(kv_a_proj_name)
@@ -977,8 +988,12 @@ class LongcatFlashForCausalLM(nn.Module):
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
futures.append(
executor.submit(weight_loader, param, loaded_weight)
maybe_executor_submit(
executor=executor,
futures=futures,
use_async=use_async_loading,
func=weight_loader,
func_args=(param, loaded_weight),
)
# Wait for all tasks to complete and raise any exceptions.