Team Ai
Apppublic

hugging-apps/padoc-document-parser

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
preprocess.py143 linesDownload Raw Back to padoc
1"""Create a PaDoc-ready checkpoint from a standard image-text model."""2 3from __future__ import annotations4 5import argparse6import json7import logging8from pathlib import Path9 10import torch11from transformers import AutoModelForImageTextToText, AutoProcessor12 13from .constants import (14    DEFAULT_FORK_TOKEN_MAP,15    DEFAULT_SPECIAL_TOKENS,16    PADOC_CONFIG_KEY,17    PADOC_FORK_MAP_KEY,18    PADOC_SPECIAL_TOKENS_KEY,19)20 21logger = logging.getLogger(__name__)22 23 24def get_padoc_metadata(model_or_config) -> dict:25    config = getattr(model_or_config, "config", model_or_config)26    metadata = getattr(config, PADOC_CONFIG_KEY, None)27    if metadata is None and hasattr(config, "text_config"):28        metadata = getattr(config.text_config, PADOC_CONFIG_KEY, None)29    if not isinstance(metadata, dict) or not metadata.get(PADOC_FORK_MAP_KEY):30        raise ValueError("Checkpoint has no padoc.fork_token_map metadata.")31    return metadata32 33 34def get_fork_token_map(model_or_config) -> dict[str, str]:35    return dict(get_padoc_metadata(model_or_config)[PADOC_FORK_MAP_KEY])36 37 38def _initialize_new_rows(model, token_ids: list[int], old_vocab_size: int, seed: int) -> None:39    if not token_ids:40        return41    with torch.no_grad(), torch.random.fork_rng():42        torch.manual_seed(seed)43        input_weights = model.get_input_embeddings().weight44        old_input = input_weights[:old_vocab_size].float()45        input_mean = old_input.mean(0)46        input_std = old_input.std(0)47        for token_id in token_ids:48            row = input_mean + torch.randn_like(input_mean) * input_std49            input_weights[token_id].copy_(row.to(input_weights.dtype))50 51        output = model.get_output_embeddings()52        if output is not None and output.weight is not input_weights:53            output_weights = output.weight54            old_output = output_weights[:old_vocab_size].float()55            output_mean = old_output.mean(0)56            output_std = old_output.std(0)57            for token_id in token_ids:58                row = output_mean + torch.randn_like(output_mean) * output_std59                output_weights[token_id].copy_(row.to(output_weights.dtype))60 61 62def preprocess_model(63    base_model: str | Path,64    output_dir: str | Path,65    *,66    special_tokens: list[str] | None = None,67    fork_token_map: dict[str, str] | None = None,68    dtype: torch.dtype = torch.bfloat16,69    seed: int = 42,70) -> Path:71    """Register atomic fork tokens and persist their mapping in config.json."""72 73    special_tokens = list(special_tokens or DEFAULT_SPECIAL_TOKENS)74    fork_token_map = dict(fork_token_map or DEFAULT_FORK_TOKEN_MAP)75    referenced = set(fork_token_map) | set(fork_token_map.values())76    if not referenced <= set(special_tokens):77        missing = sorted(referenced - set(special_tokens))78        raise ValueError(f"Fork map references tokens absent from special_tokens: {missing}")79 80    model = AutoModelForImageTextToText.from_pretrained(str(base_model), dtype=dtype)81    processor = AutoProcessor.from_pretrained(str(base_model))82    tokenizer = processor.tokenizer83    old_vocab_size = len(tokenizer)84 85    new_tokens = [86        token87        for token in special_tokens88        if len(tokenizer.encode(token, add_special_tokens=False)) != 189    ]90    if new_tokens:91        tokenizer.add_special_tokens({"additional_special_tokens": new_tokens})92        model.resize_token_embeddings(len(tokenizer))93        new_ids = [tokenizer.encode(token, add_special_tokens=False)[0] for token in new_tokens]94        _initialize_new_rows(model, new_ids, old_vocab_size, seed)95 96    for token in special_tokens:97        ids = tokenizer.encode(token, add_special_tokens=False)98        if len(ids) != 1:99            raise ValueError(f"Special token {token!r} is not atomic: {ids}")100 101    metadata = {102        PADOC_SPECIAL_TOKENS_KEY: special_tokens,103        PADOC_FORK_MAP_KEY: fork_token_map,104    }105    setattr(model.config, PADOC_CONFIG_KEY, metadata)106    if hasattr(model.config, "text_config"):107        setattr(model.config.text_config, PADOC_CONFIG_KEY, metadata)108 109    output_path = Path(output_dir).expanduser().resolve()110    output_path.mkdir(parents=True, exist_ok=True)111    model.save_pretrained(output_path)112    processor.save_pretrained(output_path)113 114    config_path = output_path / "config.json"115    with config_path.open(encoding="utf-8") as handle:116        config = json.load(handle)117    config[PADOC_CONFIG_KEY] = metadata118    with config_path.open("w", encoding="utf-8") as handle:119        json.dump(config, handle, indent=2, ensure_ascii=False)120        handle.write("\n")121    logger.info("Saved PaDoc-ready checkpoint to %s", output_path)122    return output_path123 124 125def main(argv: list[str] | None = None) -> None:126    parser = argparse.ArgumentParser(description="Create a PaDoc-ready checkpoint.")127    parser.add_argument("--base-model", required=True)128    parser.add_argument("--output", required=True)129    parser.add_argument("--seed", type=int, default=42)130    parser.add_argument("--dtype", choices=("bfloat16", "float32"), default="bfloat16")131    args = parser.parse_args(argv)132    logging.basicConfig(level=logging.INFO)133    preprocess_model(134        args.base_model,135        args.output,136        seed=args.seed,137        dtype=torch.bfloat16 if args.dtype == "bfloat16" else torch.float32,138    )139 140 141if __name__ == "__main__":142    main()143