hugging-apps/padoc-document-parser
0
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 