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 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: