From 93055059cc4dba7af3cb7b803a3f645ea3da5764 Mon Sep 17 00:00:00 2001 From: Xuan Son Nguyen Date: Tue, 11 Aug 2026 00:19:14 +0200 Subject: [PATCH] less invasive base.py --- conversion/base.py | 70 +++++++++++++++-------------------------- conversion/pockettts.py | 49 ++++++++++++++++++++++++++--- 2 files changed, 70 insertions(+), 49 deletions(-) diff --git a/conversion/base.py b/conversion/base.py index eb5a1d32a2..3572b77c21 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -58,6 +58,11 @@ logger = logging.getLogger("hf-to-gguf") AnyModel = TypeVar("AnyModel", bound="type[ModelBase]") +# for checkpoints that ship no config.json, we will try to provide a synthetic one +HparamsMatcher = Callable[[Path], bool] +HparamsLoader = Callable[[Path], dict[str, Any]] + + class SentencePieceTokenTypes(IntEnum): NORMAL = 1 UNKNOWN = 2 @@ -77,6 +82,7 @@ class ModelBase: ModelType.TEXT: {}, ModelType.MMPROJ: {}, } + _hparams_loaders: list[tuple[HparamsMatcher, HparamsLoader]] = [] dir_model: Path ftype: gguf.LlamaFileType @@ -1040,6 +1046,24 @@ class ModelBase: return part_names + @staticmethod + def load_hparams_guess(dir_model: Path) -> dict[str, Any] | None: + # some models ship no config.json, will try to guess them + from conversion import load_all_models + load_all_models() + + for matcher, loader in ModelBase._hparams_loaders: + if matcher(dir_model): + return loader(dir_model) + return None + + @classmethod + def register_hparams_loader(cls, matcher: HparamsMatcher) -> Callable[[HparamsLoader], HparamsLoader]: + def inner(loader: HparamsLoader) -> HparamsLoader: + cls._hparams_loaders.append((matcher, loader)) + return loader + return inner + @staticmethod def load_hparams(dir_model: Path, is_mistral_format: bool): if is_mistral_format: @@ -1054,7 +1078,7 @@ class ModelBase: except Exception as e: logger.warning(f"Failed to load model config from {dir_model}: {e}") if not (dir_model / "config.json").is_file(): - config = load_hparams_non_hf(dir_model) + config = ModelBase.load_hparams_guess(dir_model) if config is not None: return config logger.warning("Trying to load config.json instead") @@ -2622,50 +2646,6 @@ else: LazyTorchTensor._dtype_str_map["F8_E8M0"] = torch.uint8 -def load_hparams_non_hf(dir_model: Path) -> dict[str, Any] | None: - # some models ship no config.json at all, their hparams are derived from the checkpoint - part_names = ModelBase.get_model_part_names(dir_model, "model", ".safetensors") - if len(part_names) != 1: - return None - with gguf.utility.SafetensorsLocal(dir_model / part_names[0]) as part: - shapes = {name: tuple(part[name].shape) for name in part.keys()} - - if "flow_lm.bos_emb" in shapes: - return _load_hparams_pockettts(shapes) - - return None - - -def _load_hparams_pockettts(shapes: dict[str, tuple[int, ...]]) -> dict[str, Any]: - logger.info("gguf: detected pocket-tts checkpoint, deriving hparams from tensor shapes") - n_vocab, n_embd = shapes["flow_lm.conditioner.embed.weight"] - n_layer = sum(1 for name in shapes if re.fullmatch(r"flow_lm\.transformer\.layers\.\d+\.norm1\.weight", name)) - n_layer_a = sum(1 for name in shapes if re.fullmatch(r"mimi\.encoder_transformer\.transformer\.layers\.\d+\.norm1\.weight", name)) - n_embd_a = shapes["mimi.encoder_transformer.transformer.layers.0.norm1.weight"][0] - return { - "architectures": ["PocketTTSModel"], - "model_type": "pockettts", - "num_hidden_layers": n_layer, - "hidden_size": n_embd, - "intermediate_size": shapes["flow_lm.transformer.layers.0.linear1.weight"][0], - # the transformer is fully causal with no context limit, this only bounds the KV cache - "max_position_embeddings": 4096, - # not stored anywhere in the checkpoint, but every released variant uses head_dim 64 - "num_attention_heads": n_embd // 64, - # extra rows for the learned input vectors, see pockettts.py - # bos_before_voice only exists when the pack inserts it - "vocab_size": n_vocab + (2 if "flow_lm.bos_before_voice" in shapes else 1), - "rope_theta": 10000.0, - "layer_norm_eps": 1e-5, - "audio_config": { - "num_hidden_layers": n_layer_a, - "hidden_size": n_embd_a, - "intermediate_size": shapes["mimi.encoder_transformer.transformer.layers.0.linear1.weight"][0], - "num_attention_heads": n_embd_a // 64, - }, - } - - def get_model_architecture(hparams: dict[str, Any], model_type: ModelType) -> str: # TODO @ngxson : this won't work correctly if the model has both audio & vision encoders # maybe we should fallback to text model's arch in that case, since not many models have both diff --git a/conversion/pockettts.py b/conversion/pockettts.py index 90e16cd2b4..62ecb5acde 100644 --- a/conversion/pockettts.py +++ b/conversion/pockettts.py @@ -1,17 +1,19 @@ from __future__ import annotations -from typing import Iterable, TYPE_CHECKING +import re +from pathlib import Path +from typing import Any, Iterable, TYPE_CHECKING import torch if TYPE_CHECKING: from torch import Tensor -from .base import ModelBase, MmprojModel, SentencePieceTokenTypes, TextModel, gguf +from .base import ModelBase, MmprojModel, SentencePieceTokenTypes, TextModel, gguf, logger # Pocket TTS is a CALM: the backbone conditions a flow-matching decoder that generates one # continuous 32-d latent per frame. There is no codebook in this model. -# The checkpoint ships no config.json, hparams are derived in base.load_hparams_non_hf(). +# The checkpoint ships no config.json, hparams come from _load_hparams() below. # # Tricks being used to support this model via existing llama.cpp code paths: # - bos_before_voice and bos_emb are learned input vectors, not tokens @@ -35,6 +37,45 @@ _N_SEANET_STAGES = 3 _SAMPLE_RATE = 24000 +def _tensor_shapes(dir_model: Path) -> dict[str, tuple[int, ...]]: + part_names = ModelBase.get_model_part_names(dir_model, "model", ".safetensors") + if len(part_names) != 1: + return {} + with gguf.utility.SafetensorsLocal(dir_model / part_names[0]) as part: + return {name: tuple(part[name].shape) for name in part.keys()} + + +@ModelBase.register_hparams_loader(lambda dir_model: "flow_lm.bos_emb" in _tensor_shapes(dir_model)) +def _load_hparams(dir_model: Path) -> dict[str, Any]: + logger.info("gguf: detected pocket-tts checkpoint, deriving hparams from tensor shapes") + shapes = _tensor_shapes(dir_model) + n_vocab, n_embd = shapes["flow_lm.conditioner.embed.weight"] + n_layer = sum(1 for name in shapes if re.fullmatch(r"flow_lm\.transformer\.layers\.\d+\.norm1\.weight", name)) + n_layer_a = sum(1 for name in shapes if re.fullmatch(r"mimi\.encoder_transformer\.transformer\.layers\.\d+\.norm1\.weight", name)) + n_embd_a = shapes["mimi.encoder_transformer.transformer.layers.0.norm1.weight"][0] + return { + "architectures": ["PocketTTSModel"], + "model_type": "pockettts", + "num_hidden_layers": n_layer, + "hidden_size": n_embd, + "intermediate_size": shapes["flow_lm.transformer.layers.0.linear1.weight"][0], + # the transformer is fully causal with no context limit, this only bounds the KV cache + "max_position_embeddings": 4096, + # not in the checkpoint, but every released variant uses head_dim 64 + "num_attention_heads": n_embd // 64, + # extra rows for the learned input vectors, see _embd_table() + "vocab_size": n_vocab + (2 if "flow_lm.bos_before_voice" in shapes else 1), + "rope_theta": 10000.0, + "layer_norm_eps": 1e-5, + "audio_config": { + "num_hidden_layers": n_layer_a, + "hidden_size": n_embd_a, + "intermediate_size": shapes["mimi.encoder_transformer.transformer.layers.0.linear1.weight"][0], + "num_attention_heads": n_embd_a // 64, + }, + } + + @ModelBase.register("PocketTTSModel") class PocketTTSModel(TextModel): model_arch = gguf.MODEL_ARCH.POCKETTTS @@ -324,7 +365,7 @@ class PocketTTSMmprojModel(MmprojModel): return for stage in range(_N_SEANET_STAGES): - res_idx = _DEC_RES_IDX(stage) if is_decoder else _ENC_RES_IDX(stage) + res_idx = _DEC_RES_IDX(stage) if is_decoder else _ENC_RES_IDX(stage) scale_idx = _DEC_SCALE_IDX(stage) if is_decoder else _ENC_SCALE_IDX(stage) if idx == scale_idx: yield (self.format_tensor_name(scale, stage, suffix=suffix), data_torch)