orukeet / transformers /convert.py
NathanRoll's picture
Add qualified Transformers FP32 export for standard Parakeet clients
10b603f verified
Raw History Blame Contribute Delete
3.4 kB
"""Strict, reproducible FP32 export using the upstream Parakeet converter."""
from pathlib import Path
import hashlib
import importlib.util
import json
import tarfile
import torch
import yaml
from transformers import ParakeetForTDT, AutoProcessor
import argparse
from huggingface_hub import hf_hub_download
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("output", type=Path)
parser.add_argument("--cache-dir")
args = parser.parse_args()
ROOT = args.output.resolve()
ROOT.mkdir(parents=True, exist_ok=True)
REVISION = "555136b50265a132d4cea0d35560c26fc4f657ab"
SOURCE_SHA = "031c8ddab4845aeced904a7cde8e8aa57993b2e344716cf83a545b079c473b56"
CONVERTER_REV = "6c6bac29f50c8aad5d1f06c72b2f892a6e335dd0"
source = Path(hf_hub_download("oruk/orukeet", "orukeet-v0.1.0.nemo", revision=REVISION, cache_dir=args.cache_dir))
with source.open("rb") as handle:
assert hashlib.file_digest(handle, "sha256").hexdigest() == SOURCE_SHA
extracted = ROOT / ".nemo-source"
extracted.mkdir(exist_ok=True)
with tarfile.open(source) as archive:
archive.extractall(extracted, filter="data")
config = yaml.safe_load((extracted / "model_config.yaml").read_text())
spec = importlib.util.spec_from_file_location("upstream_converter", Path(__file__).with_name("convert_nemo_to_hf.py"))
converter = importlib.util.module_from_spec(spec)
spec.loader.exec_module(converter)
files = {"model_weights": str(extracted / "model_weights.ckpt"),
"tokenizer_model_file": str(extracted / config["tokenizer"]["model_path"].removeprefix("nemo:"))}
output = ROOT
output.mkdir(exist_ok=True)
converter.write_processor(config, files, str(output), "tdt")
model_config = converter.convert_tdt_config(config, converter.convert_encoder_config(config))
state_dict = converter.load_and_convert_tdt_state_dict(files, model_config.vocab_size)
with torch.device("meta"):
model = ParakeetForTDT(model_config)
result = model.load_state_dict(state_dict, strict=True, assign=True)
assert not result.missing_keys and not result.unexpected_keys
model.eval()
model.generation_config.decoder_start_token_id = model.config.blank_token_id
model.generation_config.suppress_tokens = list(range(model.config.vocab_size, model.config.vocab_size + len(model.config.durations)))
model.save_pretrained(output, max_shard_size="4GB")
processor = AutoProcessor.from_pretrained(output, local_files_only=True)
assert processor.tokenizer.convert_tokens_to_ids("<blank>") == model.config.blank_token_id
assert len(processor.tokenizer) == model.config.vocab_size
proof = {
"source_repo": "oruk/orukeet", "source_revision": REVISION,
"source_filename": source.name, "source_sha256": SOURCE_SHA,
"converter_url": f"https://github.com/huggingface/transformers/blob/{CONVERTER_REV}/src/transformers/models/parakeet/convert_nemo_to_hf.py",
"dtype": str(next(model.parameters()).dtype), "strict_state_dict": True,
"tensor_count": len(state_dict), "parameter_count": sum(p.numel() for p in model.parameters()),
"vocab_size": model.config.vocab_size,
"files": {}
}
for path in sorted(output.iterdir()):
if path.is_file():
with path.open("rb") as handle:
proof["files"][path.name] = {"sha256": hashlib.file_digest(handle, "sha256").hexdigest(), "bytes": path.stat().st_size}
(ROOT / "export-provenance.json").write_text(json.dumps(proof, indent=2) + "\n")
print(json.dumps(proof, indent=2))