nos-mt-es-arg-onnx / ct2_to_pegasus_dual.py
Jarbas's picture
Upload folder using huggingface_hub
69e5b1f verified
Raw
History Blame Contribute Delete
8.93 kB
"""Reconstruct a HuggingFace PegasusForConditionalGeneration checkpoint from a
CTranslate2 `model.bin` that has SEPARATE source and target vocabularies.
This is the OpenNMT-py 3.x family (Proxecto Nos). It differs from the shared
sentencepiece OpenNMT-tf models handled by ct2_to_pegasus.py:
* `encoder/embeddings_0/weight` and `decoder/embeddings/weight` are different
tables of different sizes. Pegasus has one vocabulary, so the two tables are
concatenated as [target | source]; encoder input ids are offset by the size
of the target vocabulary and the source half is made unreachable at decode
time with `final_logits_bias = -1e9`;
* position encodings ARE stored in the binary
(`decoder/position_encodings/encodings`), so no formula has to be guessed --
but Pegasus still refuses to save the table, so `_keys_to_ignore_on_save`
must be cleared;
* attention and feed-forward projections carry no bias (`add_qkvbias=False`),
so zeros are written where HuggingFace insists on one.
Usage:
python ct2_to_pegasus_dual.py <ct2_model_dir> <output_hf_dir>
"""
import json
import os
import sys
import numpy as np
import torch
from transformers import PegasusConfig, PegasusForConditionalGeneration
from ct2_reader import read_ct2_model, dequantize
from ct2_to_pegasus import ACTIVATIONS, scalar
NEG = -1e9
MAX_POS = 1024
def inspect_spec(v):
def nlayers(prefix):
return len({int(k.split("/")[1][6:]) for k in v
if k.startswith(prefix + "/layer_") and k.split("/")[1][6:].isdigit()})
tgt_size, d_model = v["decoder/embeddings/weight"].shape
src_size = v["encoder/embeddings_0/weight"].shape[0]
return dict(
encoder_layers=nlayers("encoder"),
decoder_layers=nlayers("decoder"),
source_vocab_size=int(src_size),
target_vocab_size=int(tgt_size),
d_model=int(d_model),
heads=int(scalar(v, ("num_heads", "encoder/num_heads"))),
ffn_dim=int(v["encoder/layer_0/ffn/linear_0/weight"].shape[0]),
pre_norm=bool(scalar(v, ("pre_norm", "encoder/pre_norm"))),
activation=ACTIVATIONS[int(scalar(v, ("activation", "encoder/activation")))],
layernorm_embedding=bool(scalar(v, ("layernorm_embedding",
"encoder/layernorm_embedding"), False)),
relative_position=any(k.endswith("relative_attention_bias") for k in v),
scale_embeddings=bool(scalar(v, ("decoder/scale_embeddings",))),
output_bias="decoder/projection/bias" in v,
stored_positions="decoder/position_encodings/encodings" in v,
attention_bias="encoder/layer_0/self_attention/linear_0/bias" in v,
)
def build_state_dict(v, spec):
g = lambda n: torch.from_numpy(dequantize(v, n).copy())
sd = {}
d = spec["d_model"]
n_tgt = spec["target_vocab_size"]
zero = lambda n: torch.zeros(n)
tgt_emb = g("decoder/embeddings/weight")
src_emb = g("encoder/embeddings_0/weight")
shared = torch.cat([tgt_emb, src_emb], dim=0)
sd["model.shared.weight"] = shared
sd["model.encoder.embed_tokens.weight"] = shared.clone()
sd["model.decoder.embed_tokens.weight"] = shared.clone()
# decoder/projection/weight is aliased to decoder/embeddings/weight
sd["lm_head.weight"] = shared.clone()
bias = torch.full((1, shared.shape[0]), NEG)
bias[0, :n_tgt] = g("decoder/projection/bias") if spec["output_bias"] else 0.0
sd["final_logits_bias"] = bias
pos = g("decoder/position_encodings/encodings")[:MAX_POS]
sd["model.encoder.embed_positions.weight"] = pos.clone()
sd["model.decoder.embed_positions.weight"] = pos.clone()
def ln(dst, src):
sd[dst + ".weight"] = g(src + "/gamma")
sd[dst + ".bias"] = g(src + "/beta")
def lin(dst, src, weight=None, bias=None, n=None):
w = g(src + "/weight") if weight is None else weight
sd[dst + ".weight"] = w
if bias is None:
bias = g(src + "/bias") if src and (src + "/bias") in v else zero(w.shape[0])
sd[dst + ".bias"] = bias
ln("model.encoder.layer_norm", "encoder/layer_norm")
ln("model.decoder.layer_norm", "decoder/layer_norm")
def split_bias(src, lo, hi):
return g(src + "/bias")[lo:hi] if (src + "/bias") in v else zero(hi - lo)
def self_attn(dst, src):
w = g(src + "/linear_0/weight")
for i, part in enumerate(("q_proj", "k_proj", "v_proj")):
lo, hi = i * d, (i + 1) * d
lin(dst + "." + part, None, w[lo:hi], split_bias(src + "/linear_0", lo, hi))
lin(dst + ".out_proj", src + "/linear_1")
for i in range(spec["encoder_layers"]):
s, p = "encoder/layer_%d" % i, "model.encoder.layers.%d" % i
self_attn(p + ".self_attn", s + "/self_attention")
ln(p + ".self_attn_layer_norm", s + "/self_attention/layer_norm")
lin(p + ".fc1", s + "/ffn/linear_0")
lin(p + ".fc2", s + "/ffn/linear_1")
ln(p + ".final_layer_norm", s + "/ffn/layer_norm")
for i in range(spec["decoder_layers"]):
s, p = "decoder/layer_%d" % i, "model.decoder.layers.%d" % i
self_attn(p + ".self_attn", s + "/self_attention")
ln(p + ".self_attn_layer_norm", s + "/self_attention/layer_norm")
a = s + "/attention"
lin(p + ".encoder_attn.q_proj", a + "/linear_0")
kv = g(a + "/linear_1/weight")
lin(p + ".encoder_attn.k_proj", None, kv[:d], split_bias(a + "/linear_1", 0, d))
lin(p + ".encoder_attn.v_proj", None, kv[d:], split_bias(a + "/linear_1", d, 2 * d))
lin(p + ".encoder_attn.out_proj", a + "/linear_2")
ln(p + ".encoder_attn_layer_norm", a + "/layer_norm")
lin(p + ".fc1", s + "/ffn/linear_0")
lin(p + ".fc2", s + "/ffn/linear_1")
ln(p + ".final_layer_norm", s + "/ffn/layer_norm")
return sd
def convert(ct2_dir, out_dir, src_vocab, tgt_vocab, ct2_cfg=None):
os.makedirs(out_dir, exist_ok=True)
model = read_ct2_model(os.path.join(ct2_dir, "model.bin"))
v = model["variables"]
spec = inspect_spec(v)
spec["ct2_spec"] = "%s rev %d, binary_version %d" % (
model["spec"], model["revision"], model["binary_version"])
cfg_json = ct2_cfg or {}
spec["source_eos"] = bool(cfg_json.get("add_source_eos", False))
spec["source_bos"] = bool(cfg_json.get("add_source_bos", False))
spec["decoder_start_token"] = cfg_json.get("decoder_start_token", "<s>")
print("CT2 spec:", json.dumps(spec, indent=2))
assert spec["pre_norm"] and not spec["relative_position"]
assert not spec["layernorm_embedding"]
assert spec["stored_positions"]
assert len(src_vocab) == spec["source_vocab_size"]
assert len(tgt_vocab) == spec["target_vocab_size"]
n_tgt = spec["target_vocab_size"]
cfg = PegasusConfig(
vocab_size=n_tgt + spec["source_vocab_size"], d_model=spec["d_model"],
encoder_layers=spec["encoder_layers"], decoder_layers=spec["decoder_layers"],
encoder_attention_heads=spec["heads"], decoder_attention_heads=spec["heads"],
encoder_ffn_dim=spec["ffn_dim"], decoder_ffn_dim=spec["ffn_dim"],
max_position_embeddings=MAX_POS,
activation_function=spec["activation"],
scale_embedding=spec["scale_embeddings"],
dropout=0.0, attention_dropout=0.0, activation_dropout=0.0,
pad_token_id=tgt_vocab.index("<blank>"),
bos_token_id=tgt_vocab.index("<s>"),
eos_token_id=tgt_vocab.index("</s>"),
decoder_start_token_id=tgt_vocab.index(spec["decoder_start_token"]),
forced_eos_token_id=None, max_length=512, num_beams=4,
tie_word_embeddings=True, static_position_embeddings=True,
)
hf = PegasusForConditionalGeneration(cfg)
sd = build_state_dict(v, spec)
missing, unexpected = hf.load_state_dict(sd, strict=False)
print("missing:", missing, "unexpected:", unexpected)
assert not missing and not unexpected
hf.eval()
for cls in (PegasusForConditionalGeneration,):
cls._keys_to_ignore_on_save = None
hf.save_pretrained(out_dir, safe_serialization=True)
check = PegasusForConditionalGeneration.from_pretrained(out_dir)
err = (check.model.encoder.embed_positions.weight
- hf.model.encoder.embed_positions.weight).abs().max().item()
print("position table round-trip max abs err:", err)
assert err < 1e-6, "position table was not persisted"
json.dump({"source_vocab": src_vocab, "target_vocab": tgt_vocab,
"source_offset": n_tgt, "target_vocab_size": n_tgt},
open(os.path.join(out_dir, "nos_vocab.json"), "w"), ensure_ascii=False)
return spec
if __name__ == "__main__":
d = sys.argv[1]
sv = json.load(open(os.path.join(d, "source_vocabulary.json"), encoding="utf-8"))
tv = json.load(open(os.path.join(d, "target_vocabulary.json"), encoding="utf-8"))
convert(d, sys.argv[2], sv, tv)