From 944e31fe47cf57cfc9620e1eb0261db0bf3f77f3 Mon Sep 17 00:00:00 2001 From: Masahiro Hiramori Date: Fri, 7 Aug 2026 07:58:04 +0000 Subject: [PATCH 1/2] avoid CUDA OOM when loading Moshi checkpoints --- src/kame/models/loaders.py | 24 +++++++++++++----------- 1 file changed, 13 insertions(+), 11 deletions(-) diff --git a/src/kame/models/loaders.py b/src/kame/models/loaders.py index 3df8f75..8c092de 100644 --- a/src/kame/models/loaders.py +++ b/src/kame/models/loaders.py @@ -370,16 +370,18 @@ def _upgrade_legacy_lm_state_dict(state: dict[str, torch.Tensor]) -> dict[str, t return upgraded_state -def _cast_lm_state_dict(state: dict[str, torch.Tensor], dtype: torch.dtype) -> dict[str, torch.Tensor]: - """Cast LM checkpoint tensors to the runtime dtype expected by the current model.""" - state = dict(state) +def _cast_lm_state_dict( + state: dict[str, torch.Tensor], + dtype: torch.dtype, + device: torch.device | str, +) -> dict[str, torch.Tensor]: + """Move LM checkpoint tensors to the runtime device and expected dtype.""" for key, value in state.items(): if value.dtype.is_floating_point: - if key.startswith("condition_provider.") or key.startswith("fuser."): - value = value.float() - else: - value = value.to(dtype) - state[key] = value + target_dtype = torch.float32 if key.startswith(("condition_provider.", "fuser.")) else dtype + state[key] = value.to(device=device, dtype=target_dtype) + else: + state[key] = value.to(device=device) return state @@ -427,15 +429,15 @@ def get_moshi_lm( if filename is not None: if _is_safetensors(filename): - state = load_file(filename, device=str(device)) + state = load_file(filename, device="cpu") else: pkg = torch.load( filename, - map_location=device, + map_location="cpu", ) state = pkg["fsdp_best_state"]["model"] state = _upgrade_legacy_lm_state_dict(state) - state = _cast_lm_state_dict(state, dtype) + state = _cast_lm_state_dict(state, dtype, device) model.load_state_dict(state, assign=True) if lora: From 37ec217616a8ebd85ba162b89adbe644d3254cce Mon Sep 17 00:00:00 2001 From: So Kuroki Date: Fri, 7 Aug 2026 11:43:06 +0000 Subject: [PATCH 2/2] Fix Ruff 0.16 lint configuration --- pyproject.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index f2a115a..2c4c074 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -73,5 +73,8 @@ quantization = [ [tool.ruff] line-length = 120 +[tool.ruff.lint] +select = ["E4", "E7", "E9", "F"] + [tool.pyright] reportPrivateImportUsage = false