[diffusion] chore: optimize model weight loading (#34064)

This commit is contained in:
Mick
2026-08-10 20:07:48 +08:00
committed by GitHub
parent 4eaaeda004
commit 8ba9385097
8 changed files with 299 additions and 28 deletions
+1
View File
@@ -81,6 +81,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
- `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks. - `--lora-merge-mode {auto|merge|dynamic}`: choose how LoRA is applied. `auto` statically merges regular weights and uses dynamic LoRA for FSDP-sharded weights to avoid full-gather peaks.
- `--num-gpus {N}`: number of GPUs to use - `--num-gpus {N}`: number of GPUs to use
- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes. - `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths. `speed` keeps `torch.compile` disabled unless a model-specific deployment config opts in after validation; pass `--enable-torch-compile true` to enable it explicitly. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes.
- `--direct-gpu-weight-loading {true|false}`: opt into direct GPU loading for an unquantized, GPU-resident, TP=1 DiT by materializing its complete checkpoint state dict on GPU. Startup impact is model-dependent, so benchmark the target model before deployment. Disabled by default because checkpoint and model weights coexist temporarily, substantially increasing peak GPU memory. It is incompatible with DiT CPU/layerwise offload and FSDP.
- `--tp-size {N}`: tensor parallelism size. Depending on the pipeline, it can shard the DiT, one or more encoders, or both. - `--tp-size {N}`: tensor parallelism size. Depending on the pipeline, it can shard the DiT, one or more encoders, or both.
- `--sp-degree {N}`: sequence parallelism size - `--sp-degree {N}`: sequence parallelism size
- `--dp-size {N}` (alias `--data-parallel-size`): number of data-parallel replicas. Each replica is a full copy of the engine on `num_gpus / N` GPUs with its own ingress; generation requests round-robin across replicas, realtime sessions stick to the replica holding their state, and control operations (weights, LoRA, memory occupation, shutdown) apply to every replica. Combines with the other parallelism axes (`num_gpus = dp × cfg × tp × sp`); monolithic serving only. - `--dp-size {N}` (alias `--data-parallel-size`): number of data-parallel replicas. Each replica is a full copy of the engine on `num_gpus / N` GPUs with its own ingress; generation requests round-robin across replicas, realtime sessions stick to the replica holding their state, and control operations (weights, LoRA, memory occupation, shutdown) apply to every replica. Combines with the other parallelism axes (`num_gpus = dp × cfg × tp × sp`); monolithic serving only.
@@ -39,6 +39,17 @@ _is_npu = is_npu()
logger = init_logger(__name__) logger = init_logger(__name__)
def _resolve_checkpoint_load_device(
runtime_device: torch.device,
*,
component_cpu_offload: bool,
runtime_quant_config: object | None,
) -> torch.device:
if component_cpu_offload and runtime_quant_config is None:
return torch.device("cpu")
return runtime_device
def _default_quantized_attention_backend( def _default_quantized_attention_backend(
quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs
) -> AttentionBackendEnum | None: ) -> AttentionBackendEnum | None:
@@ -198,10 +209,29 @@ class TransformerLoader(ComponentLoader):
logger.debug("quantization config: %s", init_params["quant_config"]) logger.debug("quantization config: %s", init_params["quant_config"])
local_torch_device = get_local_torch_device() local_torch_device = get_local_torch_device()
checkpoint_load_device = _resolve_checkpoint_load_device(
local_torch_device,
component_cpu_offload=bool(component_server_args.dit_cpu_offload),
runtime_quant_config=quant_spec.runtime_quant_config,
)
direct_gpu_weight_loading = bool(
component_server_args.direct_gpu_weight_loading
)
if direct_gpu_weight_loading and quant_spec.runtime_quant_config is not None:
raise ValueError(
"--direct-gpu-weight-loading supports only unquantized DiT checkpoints"
)
weight_load_plan = WeightLoadPlan.for_component( weight_load_plan = WeightLoadPlan.for_component(
checkpoint_load_device=local_torch_device, checkpoint_load_device=checkpoint_load_device,
needs_device_weight_postprocess=quant_spec.needs_device_weight_postprocess, needs_device_weight_postprocess=quant_spec.needs_device_weight_postprocess,
component_cpu_offload=bool(component_server_args.dit_cpu_offload), component_cpu_offload=bool(component_server_args.dit_cpu_offload),
load_full_state_dict_on_device=direct_gpu_weight_loading,
)
if direct_gpu_weight_loading:
logger.warning(
"Direct GPU weight loading is enabled for %s; the complete checkpoint "
"state dict and materialized model weights may coexist on GPU during startup",
component_name,
) )
quantized_attn_backend = _default_quantized_attention_backend( quantized_attn_backend = _default_quantized_attention_backend(
@@ -9,6 +9,7 @@
from collections import Counter, defaultdict from collections import Counter, defaultdict
from collections.abc import Callable, Generator from collections.abc import Callable, Generator
from itertools import chain from itertools import chain
from types import MethodType
from typing import Any from typing import Any
import torch import torch
@@ -26,7 +27,12 @@ from torch.distributed.fsdp import (
from torch.nn.modules.module import _IncompatibleKeys from torch.nn.modules.module import _IncompatibleKeys
from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in from sglang.multimodal_gen.configs.models.fsdp import is_module_list_entry_in
from sglang.multimodal_gen.runtime.layers.linear import UnquantizedLinearMethod from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
UnquantizedLinearMethod,
)
from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import ( from sglang.multimodal_gen.runtime.layers.quantization.bitsandbytes import (
attach_bitsandbytes_4bit_quant_states, attach_bitsandbytes_4bit_quant_states,
build_bitsandbytes_4bit_quant_states, build_bitsandbytes_4bit_quant_states,
@@ -93,6 +99,45 @@ def _make_param_like(
return new_param return new_param
def _can_assign_cpu_tensor_without_copy(
actual_param: torch.nn.Parameter,
full_tensor: torch.Tensor,
target_param: torch.Tensor,
) -> bool:
"""Return whether a TP=1 linear loader would only copy this CPU tensor."""
if full_tensor.device.type != "cpu":
return False
weight_loader = actual_param.__dict__.get("weight_loader")
if not isinstance(weight_loader, MethodType):
return False
owner = weight_loader.__self__
if not isinstance(
owner,
(ReplicatedLinear, ColumnParallelLinear, RowParallelLinear),
):
return False
if not isinstance(owner.quant_method, UnquantizedLinearMethod):
return False
if not isinstance(owner, ReplicatedLinear) and owner.tp_size != 1:
return False
if type(actual_param) is not nn.Parameter:
return False
if any(
actual_param.__dict__.get(attribute, False)
for attribute in (
"is_metadata",
"is_sharded_weight",
"needs_scalar_to_array",
)
):
return False
return (
full_tensor.shape == target_param.shape
and full_tensor.dtype == target_param.dtype
)
def _make_class_name_shard_condition(class_names: set[str]): def _make_class_name_shard_condition(class_names: set[str]):
def shard_condition(n: str, m: nn.Module) -> bool: def shard_condition(n: str, m: nn.Module) -> bool:
return type(m).__name__ in class_names return type(m).__name__ in class_names
@@ -282,7 +327,8 @@ def maybe_load_fsdp_model(
preconverted_state_dict = None preconverted_state_dict = None
is_bnb_quantized = _is_bitsandbytes_quant_config(init_params.get("quant_config")) is_bnb_quantized = _is_bitsandbytes_quant_config(init_params.get("quant_config"))
if ( if (
use_fsdp not weight_load_plan.load_full_state_dict_on_device
and use_fsdp
and weight_dir_list and weight_dir_list
and preprocess_loaded_state_dict is None and preprocess_loaded_state_dict is None
and not is_bnb_quantized and not is_bnb_quantized
@@ -295,7 +341,8 @@ def maybe_load_fsdp_model(
) )
) )
elif ( elif (
not use_fsdp not weight_load_plan.load_full_state_dict_on_device
and not use_fsdp
and weight_dir_list and weight_dir_list
and preprocess_loaded_state_dict is None and preprocess_loaded_state_dict is None
and not is_bnb_quantized and not is_bnb_quantized
@@ -309,6 +356,12 @@ def maybe_load_fsdp_model(
) )
if preconverted_state_dict is None: if preconverted_state_dict is None:
if weight_load_plan.load_full_state_dict_on_device:
weight_iterator = safetensors_weights_iterator(
weight_dir_list,
weight_load_plan=weight_load_plan,
)
else:
weight_iterator = safetensors_weights_iterator(weight_dir_list) weight_iterator = safetensors_weights_iterator(weight_dir_list)
if preprocess_loaded_state_dict is not None: if preprocess_loaded_state_dict is not None:
weight_iterator = preprocess_loaded_state_dict(weight_iterator) weight_iterator = preprocess_loaded_state_dict(weight_iterator)
@@ -611,7 +664,8 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu() sharded_tensor = sharded_tensor.cpu()
elif not isinstance(meta_sharded_param, dist_tensor.DTensor): elif not isinstance(meta_sharded_param, dist_tensor.DTensor):
full_tensor = full_tensor.to( full_tensor = full_tensor.to(
device=checkpoint_load_device, dtype=target_dtype device=checkpoint_load_device,
dtype=target_dtype,
) )
actual_param = rank_local_checkpoint.get_param_for_weight_loading( actual_param = rank_local_checkpoint.get_param_for_weight_loading(
model, param_dict, target_param_name model, param_dict, target_param_name
@@ -623,16 +677,24 @@ def load_model_from_full_model_state_dict(
) )
if weight_loader is not None: if weight_loader is not None:
assert actual_param is not None assert actual_param is not None
if _can_assign_cpu_tensor_without_copy(
actual_param,
full_tensor,
meta_sharded_param,
):
sharded_tensor = full_tensor
else:
sharded_tensor = torch.empty_like( sharded_tensor = torch.empty_like(
meta_sharded_param, meta_sharded_param,
device=checkpoint_load_device, device=checkpoint_load_device,
dtype=target_dtype, dtype=target_dtype,
) )
# Preserve requires_grad flag to avoid errors with non-floating dtypes # Preserve requires_grad flag to avoid errors with non-floating dtypes
requires_grad = getattr(meta_sharded_param, "requires_grad", False) requires_grad = meta_sharded_param.requires_grad
temp_param = _make_param_like(actual_param, sharded_tensor) temp_param = _make_param_like(actual_param, sharded_tensor)
if not ( if not (
sharded_tensor.is_floating_point() or sharded_tensor.is_complex() sharded_tensor.is_floating_point()
or sharded_tensor.is_complex()
): ):
requires_grad = False requires_grad = False
temp_param.requires_grad = requires_grad temp_param.requires_grad = requires_grad
@@ -13,6 +13,8 @@ class WeightLoadPlan:
weight_postprocess_device: torch.device | None = None weight_postprocess_device: torch.device | None = None
# Delay non-FSDP component CPU offload until after weight postprocessing. # Delay non-FSDP component CPU offload until after weight postprocessing.
defer_component_cpu_offload: bool = False defer_component_cpu_offload: bool = False
# keep the complete mapped checkpoint state dict on the load device
load_full_state_dict_on_device: bool = False
@classmethod @classmethod
def for_component( def for_component(
@@ -21,6 +23,7 @@ class WeightLoadPlan:
checkpoint_load_device: torch.device, checkpoint_load_device: torch.device,
needs_device_weight_postprocess: bool, needs_device_weight_postprocess: bool,
component_cpu_offload: bool, component_cpu_offload: bool,
load_full_state_dict_on_device: bool = False,
) -> "WeightLoadPlan": ) -> "WeightLoadPlan":
# if on-device weight postprocessing is required, load directly to device to speedup loading # if on-device weight postprocessing is required, load directly to device to speedup loading
weight_postprocess_device = ( weight_postprocess_device = (
@@ -32,4 +35,5 @@ class WeightLoadPlan:
defer_component_cpu_offload=( defer_component_cpu_offload=(
needs_device_weight_postprocess and component_cpu_offload needs_device_weight_postprocess and component_cpu_offload
), ),
load_full_state_dict_on_device=load_full_state_dict_on_device,
) )
@@ -299,6 +299,8 @@ class ServerArgs(DisaggServerArgsMixin):
# CPU offload parameters # CPU offload parameters
dit_cpu_offload: bool | None = None dit_cpu_offload: bool | None = None
# trade checkpoint-loading peak memory for faster ordinary DiT startup
direct_gpu_weight_loading: bool = False
# if true, select the DiT layerwise group # if true, select the DiT layerwise group
dit_layerwise_offload: bool | None = None dit_layerwise_offload: bool | None = None
layerwise_offload_components: list[str] | None = None layerwise_offload_components: list[str] | None = None
@@ -508,6 +510,7 @@ class ServerArgs(DisaggServerArgsMixin):
self._validate_scheduler_rpc_timeout() self._validate_scheduler_rpc_timeout()
self._validate_pipeline() self._validate_pipeline()
self._validate_offload() self._validate_offload()
self._validate_direct_gpu_weight_loading()
if not current_platform.is_cpu(): if not current_platform.is_cpu():
self._validate_parallelism() self._validate_parallelism()
self._validate_cfg_parallel() self._validate_cfg_parallel()
@@ -1771,6 +1774,15 @@ class ServerArgs(DisaggServerArgsMixin):
action=StoreBoolean, action=StoreBoolean,
help="Use CPU offload for DiT inference. Enable if run out of memory with FSDP.", help="Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
) )
parser.add_argument(
"--direct-gpu-weight-loading",
action=StoreBoolean,
default=ServerArgs.direct_gpu_weight_loading,
help="Load the full unquantized DiT checkpoint state dict directly "
"onto GPU before assigning model parameters. This may reduce startup "
"time depending on the model, but temporarily requires checkpoint "
"weights and model weights to coexist on GPU. Disabled by default.",
)
parser.add_argument( parser.add_argument(
"--dit-layerwise-offload", "--dit-layerwise-offload",
action=StoreBoolean, action=StoreBoolean,
@@ -2634,6 +2646,23 @@ class ServerArgs(DisaggServerArgsMixin):
"--performance-mode speed for GPU-resident defaults when memory allows." "--performance-mode speed for GPU-resident defaults when memory allows."
) )
def _validate_direct_gpu_weight_loading(self) -> None:
if not self.direct_gpu_weight_loading:
return
if not current_platform.is_cuda():
raise ValueError("--direct-gpu-weight-loading requires CUDA")
if self.dit_cpu_offload or self.is_dit_layerwise_offload_selected:
raise ValueError(
"--direct-gpu-weight-loading requires a GPU-resident DiT; disable "
"DiT CPU and layerwise offload"
)
if self.use_fsdp_inference:
raise ValueError(
"--direct-gpu-weight-loading does not support FSDP inference"
)
if self.tp_size != 1:
raise ValueError("--direct-gpu-weight-loading requires --tp-size 1")
def _validate_parallelism(self): def _validate_parallelism(self):
if self.kv_gather_degree < 1: if self.kv_gather_degree < 1:
raise ValueError("kv_gather_degree must be >= 1") raise ValueError("kv_gather_degree must be >= 1")
@@ -9,7 +9,9 @@ import torch
from safetensors.torch import safe_open, save_file from safetensors.torch import safe_open, save_file
from torch import nn from torch import nn
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
from sglang.multimodal_gen.runtime.loader import fsdp_load, rank_local_checkpoint from sglang.multimodal_gen.runtime.loader import fsdp_load, rank_local_checkpoint
from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan
class _UniformDtypeModel(nn.Module): class _UniformDtypeModel(nn.Module):
@@ -26,6 +28,12 @@ class _MixedDtypeModel(_UniformDtypeModel):
_fsdp_mixed_dtype_params = True _fsdp_mixed_dtype_params = True
class _ReplicatedLinearModel(_UniformDtypeModel):
def __init__(self) -> None:
super().__init__()
self.proj = ReplicatedLinear(4, 4, bias=False)
class TestFSDPMixedPrecisionPolicy(unittest.TestCase): class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
def _load_and_capture_policy( def _load_and_capture_policy(
self, self,
@@ -97,6 +105,61 @@ class TestFSDPMixedPrecisionPolicy(unittest.TestCase):
shard_model.assert_not_called() shard_model.assert_not_called()
class TestOrdinaryWeightLoading(unittest.TestCase):
def test_direct_device_loading_skips_rank_local_cpu_checkpoint(self):
load_plan = WeightLoadPlan(
checkpoint_load_device=torch.device("cuda:0"),
load_full_state_dict_on_device=True,
)
with (
patch.object(fsdp_load.current_platform, "is_mps", return_value=False),
patch.object(
rank_local_checkpoint,
"try_load_rank_local_tp_state_dict",
) as rank_local_load,
patch.object(
fsdp_load,
"safetensors_weights_iterator",
return_value=iter(()),
) as weight_iterator,
patch.object(fsdp_load, "load_model_from_full_model_state_dict"),
):
fsdp_load.maybe_load_fsdp_model(
model_cls=_UniformDtypeModel,
init_params={},
weight_dir_list=["model.safetensors"],
device=torch.device("cuda:0"),
hsdp_replicate_dim=1,
hsdp_shard_dim=1,
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
weight_load_plan=load_plan,
)
rank_local_load.assert_not_called()
weight_iterator.assert_called_once_with(
["model.safetensors"],
weight_load_plan=load_plan,
)
def test_tp1_unquantized_linear_assigns_checkpoint_tensor_without_copy(self):
with torch.device("meta"):
model = _ReplicatedLinearModel()
checkpoint_weight = torch.arange(16, dtype=torch.float32).reshape(4, 4)
fsdp_load.load_model_from_full_model_state_dict(
model,
iter((("proj.weight", checkpoint_weight),)),
checkpoint_load_device=torch.device("cpu"),
param_dtype=torch.float32,
strict=True,
param_names_mapping=fsdp_load.get_param_names_mapping({}),
)
self.assertEqual(model.proj.weight.data_ptr(), checkpoint_weight.data_ptr())
torch.testing.assert_close(model.proj.weight, checkpoint_weight)
class TestRankLocalSafetensorsRead(unittest.TestCase): class TestRankLocalSafetensorsRead(unittest.TestCase):
def _source( def _source(
self, self,
@@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import ( from sglang.multimodal_gen.runtime.pipelines.minimax_h3_pipeline import (
MiniMaxH3Pipeline, MiniMaxH3Pipeline,
) )
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ( from sglang.multimodal_gen.runtime.server_args import (
MAX_SCHEDULER_RPC_TIMEOUT_S, MAX_SCHEDULER_RPC_TIMEOUT_S,
ServerArgs, ServerArgs,
@@ -2302,5 +2303,44 @@ class TestNcclNvlsArgs(unittest.TestCase):
self.assertFalse(disabled_args.enable_nccl_nvls) self.assertFalse(disabled_args.enable_nccl_nvls)
class TestDirectGpuWeightLoading(unittest.TestCase):
def _args(self) -> ServerArgs:
args = ServerArgs.__new__(ServerArgs)
args.direct_gpu_weight_loading = True
args.dit_cpu_offload = False
args.layerwise_offload_components = []
args.use_fsdp_inference = False
args.tp_size = 1
return args
def test_cli_defaults_off_and_parses_explicit_enable(self):
parser = FlexibleArgumentParser()
ServerArgs.add_cli_args(parser)
default_args, _ = parser.parse_known_args(["--model-path", "/fake"])
enabled_args, _ = parser.parse_known_args(
["--model-path", "/fake", "--direct-gpu-weight-loading"]
)
self.assertFalse(default_args.direct_gpu_weight_loading)
self.assertTrue(enabled_args.direct_gpu_weight_loading)
def test_rejects_cpu_offload_fsdp_and_tp(self):
cpu_offload_args = self._args()
cpu_offload_args.dit_cpu_offload = True
fsdp_args = self._args()
fsdp_args.use_fsdp_inference = True
tp_args = self._args()
tp_args.tp_size = 2
with patch.object(current_platform, "is_cuda", return_value=True):
with self.assertRaisesRegex(ValueError, "GPU-resident DiT"):
cpu_offload_args._validate_direct_gpu_weight_loading()
with self.assertRaisesRegex(ValueError, "FSDP"):
fsdp_args._validate_direct_gpu_weight_loading()
with self.assertRaisesRegex(ValueError, "tp-size 1"):
tp_args._validate_direct_gpu_weight_loading()
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.layers.quantization.modelopt_quant import (
from sglang.multimodal_gen.runtime.loader.component_loaders import transformer_loader from sglang.multimodal_gen.runtime.loader.component_loaders import transformer_loader
from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader import (
_default_quantized_attention_backend, _default_quantized_attention_backend,
_resolve_checkpoint_load_device,
_warn_if_expected_param_dtype_missing, _warn_if_expected_param_dtype_missing,
) )
from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( from sglang.multimodal_gen.runtime.loader.transformer_load_utils import (
@@ -118,6 +119,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
quantization_ignored_layers=None, quantization_ignored_layers=None,
tp_size=1, tp_size=1,
dit_cpu_offload=False, dit_cpu_offload=False,
direct_gpu_weight_loading=False,
text_encoder_cpu_offload=False, text_encoder_cpu_offload=False,
) )
defaults.update(overrides) defaults.update(overrides)
@@ -253,6 +255,46 @@ class TestTransformerQuantHelpers(unittest.TestCase):
self.assertEqual(plan.checkpoint_load_device, device) self.assertEqual(plan.checkpoint_load_device, device)
self.assertEqual(plan.weight_postprocess_device, device) self.assertEqual(plan.weight_postprocess_device, device)
self.assertTrue(plan.defer_component_cpu_offload) self.assertTrue(plan.defer_component_cpu_offload)
self.assertFalse(plan.load_full_state_dict_on_device)
def test_weight_load_plan_can_keep_full_state_dict_on_device(self):
plan = WeightLoadPlan.for_component(
checkpoint_load_device=torch.device("cuda:0"),
needs_device_weight_postprocess=False,
component_cpu_offload=False,
load_full_state_dict_on_device=True,
)
self.assertTrue(plan.load_full_state_dict_on_device)
def test_unquantized_cpu_offload_loads_checkpoint_on_cpu(self):
device = _resolve_checkpoint_load_device(
torch.device("cuda:0"),
component_cpu_offload=True,
runtime_quant_config=None,
)
self.assertEqual(device, torch.device("cpu"))
def test_quantized_cpu_offload_keeps_checkpoint_on_runtime_device(self):
runtime_device = torch.device("cuda:0")
device = _resolve_checkpoint_load_device(
runtime_device,
component_cpu_offload=True,
runtime_quant_config=object(),
)
self.assertEqual(device, runtime_device)
def test_resident_transformer_loads_checkpoint_on_runtime_device(self):
runtime_device = torch.device("cuda:0")
device = _resolve_checkpoint_load_device(
runtime_device,
component_cpu_offload=False,
runtime_quant_config=None,
)
self.assertEqual(device, runtime_device)
def test_mixed_model_with_expected_dtype_does_not_warn(self): def test_mixed_model_with_expected_dtype_does_not_warn(self):
model = torch.nn.Module() model = torch.nn.Module()