Feature Extraction
Transformers
Safetensors
Korean
han2han
text-generation
hanja
hangul
historical-korean
encoder-decoder
custom_code
Instructions to use cadazar/han2han-pt with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use cadazar/han2han-pt with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="cadazar/han2han-pt", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("cadazar/han2han-pt", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Han2Han pre-trained checkpoint (uniform average of the last five checkpoints, steps 6347 to 6725)
Browse files- README.md +128 -0
- chat_template.jinja +24 -0
- config.json +111 -0
- generation_config.json +8 -0
- han2han_config.py +565 -0
- model.safetensors +3 -0
- modeling_han2han.py +0 -0
- spiece.model +3 -0
- tokenizer.json +0 -0
- tokenizer_config.json +15 -0
README.md
ADDED
|
@@ -0,0 +1,128 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: apache-2.0
|
| 3 |
+
language:
|
| 4 |
+
- ko
|
| 5 |
+
library_name: transformers
|
| 6 |
+
pipeline_tag: feature-extraction
|
| 7 |
+
tags:
|
| 8 |
+
- han2han
|
| 9 |
+
- hanja
|
| 10 |
+
- hangul
|
| 11 |
+
- historical-korean
|
| 12 |
+
- encoder-decoder
|
| 13 |
+
- custom_code
|
| 14 |
+
---
|
| 15 |
+
|
| 16 |
+
# Han2Han PT
|
| 17 |
+
|
| 18 |
+
The pre-trained (PT) checkpoint of [Han2Han](https://github.com/cadazar/han2han), a
|
| 19 |
+
169M-parameter encoder-decoder model that learns script-invariant
|
| 20 |
+
representations of Korean text: a document written in Hanja and its Hangul
|
| 21 |
+
transcription land at the same point in embedding space. The recipe (jamo and
|
| 22 |
+
character-level embedding fusion, morpheme-aware denoising, bidirectional
|
| 23 |
+
Hanja-Hangul transcription) is described in the paper, accepted to Findings of
|
| 24 |
+
EMNLP 2026.
|
| 25 |
+
|
| 26 |
+
This is the starting point for fine-tuning, and the first of three checkpoints:
|
| 27 |
+
|
| 28 |
+
| Repo | Stage |
|
| 29 |
+
| --- | --- |
|
| 30 |
+
| `cadazar/han2han-pt` (this one) | pre-training, 35B tokens |
|
| 31 |
+
| [`cadazar/han2han-it`](https://huggingface.co/cadazar/han2han-it) | instruction tuning from these weights |
|
| 32 |
+
| [`cadazar/han2han-rl`](https://huggingface.co/cadazar/han2han-rl) | reinforcement learning on Hanja-Hangul transcription from `han2han-it` |
|
| 33 |
+
|
| 34 |
+
## These weights are an average
|
| 35 |
+
|
| 36 |
+
The weights are the uniform average of the last five checkpoints of the
|
| 37 |
+
pre-training run (steps 6347, 6441, 6536, 6630 and 6725, saved every 0.5B
|
| 38 |
+
tokens from 33B to 35B), not the final checkpoint alone. The average is what
|
| 39 |
+
the instruction tuning started from, chosen after comparing starting points,
|
| 40 |
+
so it is the pre-trained model that `han2han-it` and `han2han-rl` descend
|
| 41 |
+
from. The learning rate had nearly finished its cooldown over those steps: no
|
| 42 |
+
weight differs from the final checkpoint by more than 0.0015.
|
| 43 |
+
|
| 44 |
+
## Intended use
|
| 45 |
+
|
| 46 |
+
Fine-tuning on downstream tasks, and sentence embeddings that treat Hanja and
|
| 47 |
+
Hangul spellings of the same text alike.
|
| 48 |
+
|
| 49 |
+
It is not a chat model. The tokenizer and the chat template are the same as in
|
| 50 |
+
the other two repos, but pre-training never used the chat tokens
|
| 51 |
+
(`<|system|>`, `<|user|>`, `<|assistant|>`, `<|think|>`, `<|end_of_turn|>`):
|
| 52 |
+
their embedding rows are still nearly one shared vector here, as for every
|
| 53 |
+
other token pre-training never saw. For generation from a prompt use
|
| 54 |
+
`han2han-it` or `han2han-rl`.
|
| 55 |
+
|
| 56 |
+
## Usage
|
| 57 |
+
|
| 58 |
+
Runtime requirements: `torch` and `transformers` (tested with torch 2.14.1 and
|
| 59 |
+
transformers 5.18.0 on CPU). The model class ships in this repo, so loading needs
|
| 60 |
+
`trust_remote_code=True`; the tokenizer itself runs no custom code, but without
|
| 61 |
+
the flag `AutoTokenizer` stops to ask.
|
| 62 |
+
|
| 63 |
+
`output_sentence_embeddings=True` returns the mean-pooled encoder states, one
|
| 64 |
+
640-dimensional vector per input.
|
| 65 |
+
|
| 66 |
+
```python
|
| 67 |
+
import torch
|
| 68 |
+
import torch.nn.functional as F
|
| 69 |
+
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
|
| 70 |
+
|
| 71 |
+
repo = "cadazar/han2han-pt"
|
| 72 |
+
tokenizer = AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
|
| 73 |
+
model = AutoModelForSeq2SeqLM.from_pretrained(repo, trust_remote_code=True).eval()
|
| 74 |
+
|
| 75 |
+
texts = [
|
| 76 |
+
"會場을 一巡하고 돌아올 때까지도 日本人側 畵家까지도 一人도 發見할 수가 없었다.",
|
| 77 |
+
"회장을 일순하고 돌아올 때까지도 일본인측 화가까지도 일인도 발견할 수가 없었다.",
|
| 78 |
+
"南畵나 四君子에서는 破墨의 妙法으로 白雪을 象徵할 수 있다.",
|
| 79 |
+
]
|
| 80 |
+
inputs = tokenizer(texts, return_tensors="pt", padding=True)
|
| 81 |
+
with torch.no_grad():
|
| 82 |
+
embeddings = model(**inputs, output_sentence_embeddings=True)[0]
|
| 83 |
+
embeddings = F.normalize(embeddings, dim=-1)
|
| 84 |
+
print(embeddings @ embeddings.T)
|
| 85 |
+
# tensor([[1.0000, 0.9315, 0.8253],
|
| 86 |
+
# [0.9315, 1.0000, 0.7910],
|
| 87 |
+
# [0.8253, 0.7910, 1.0000]])
|
| 88 |
+
```
|
| 89 |
+
|
| 90 |
+
The first two rows are the same 1920s newspaper sentence in mixed script and in
|
| 91 |
+
Hangul; the third is a different sentence.
|
| 92 |
+
|
| 93 |
+
### Tokenizer
|
| 94 |
+
|
| 95 |
+
`tokenizer.json` is a `tokenizers`-library build of the SentencePiece model
|
| 96 |
+
(`spiece.model`, kept here as the source), made by
|
| 97 |
+
`scripts/build_hf_tokenizer.py` in the GitHub repo. Calling the tokenizer does
|
| 98 |
+
not add BOS or EOS tokens, and special tokens written into the text are mapped
|
| 99 |
+
to their ids. It encodes like the SentencePiece wrapper the training code uses,
|
| 100 |
+
with one known difference: where two segmentations of a span have exactly the
|
| 101 |
+
same score (runs of digits, mostly), the same pieces can come out in a different
|
| 102 |
+
order.
|
| 103 |
+
|
| 104 |
+
## Files
|
| 105 |
+
|
| 106 |
+
| File | Contents |
|
| 107 |
+
| --- | --- |
|
| 108 |
+
| `model.safetensors` | fp32 weights, 169.2M parameters, plus the `jbu` / `cbu` subword bucket tables |
|
| 109 |
+
| `config.json`, `generation_config.json` | model and generation config, with `auto_map` entries for the Auto classes |
|
| 110 |
+
| `tokenizer.json`, `tokenizer_config.json`, `chat_template.jinja` | fast tokenizer (38400 pieces) and chat template, the same as in `han2han-it` |
|
| 111 |
+
| `spiece.model` | the SentencePiece model `tokenizer.json` was built from |
|
| 112 |
+
| `modeling_han2han.py`, `han2han_config.py` | modeling code from the GitHub repo at commit [`af1330e`](https://github.com/cadazar/han2han/commit/af1330e49f1edf46e26b63937724ddbaeffc3bee), the same files as in `han2han-it` |
|
| 113 |
+
|
| 114 |
+
## Citation
|
| 115 |
+
|
| 116 |
+
```bibtex
|
| 117 |
+
@inproceedings{han2han2026,
|
| 118 |
+
title = {Han2Han: Efficient Language-Specific Character Representation
|
| 119 |
+
through Script-Aware Pre-Training for Historical Text Analysis},
|
| 120 |
+
author = {Adams, Cellik and Jo, EunKyoung and Kim, Juae},
|
| 121 |
+
booktitle = {Findings of the Association for Computational Linguistics: EMNLP 2026},
|
| 122 |
+
year = {2026}
|
| 123 |
+
}
|
| 124 |
+
```
|
| 125 |
+
|
| 126 |
+
## License
|
| 127 |
+
|
| 128 |
+
Apache License 2.0, the same as the [GitHub repo](https://github.com/cadazar/han2han).
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- if messages and messages[0]['role'] == 'system' -%}
|
| 2 |
+
{%- if messages[0]['content'] is string -%}
|
| 3 |
+
{{- '<|system|>' + messages[0]['content'] -}}
|
| 4 |
+
{%- else -%}
|
| 5 |
+
{{- '<|system|>' + messages[0]['content'] | selectattr('type', 'equalto', 'text') | map(attribute='text') | join('') -}}
|
| 6 |
+
{%- endif -%}
|
| 7 |
+
{%- set loop_messages = messages[1:] -%}
|
| 8 |
+
{%- else -%}
|
| 9 |
+
{%- set loop_messages = messages -%}
|
| 10 |
+
{%- endif -%}
|
| 11 |
+
{%- for message in loop_messages -%}
|
| 12 |
+
{%- if message['role'] not in ['user', 'assistant'] -%}
|
| 13 |
+
{{- raise_exception('Han2Han chat supports a leading system message, then user and assistant turns; got role ' + message['role']) -}}
|
| 14 |
+
{%- endif -%}
|
| 15 |
+
{%- if message['content'] is string -%}
|
| 16 |
+
{%- set content = message['content'] -%}
|
| 17 |
+
{%- else -%}
|
| 18 |
+
{%- set content = message['content'] | selectattr('type', 'equalto', 'text') | map(attribute='text') | join('') -%}
|
| 19 |
+
{%- endif -%}
|
| 20 |
+
{{- '<|' + message['role'] + '|>' + content + '<|end_of_turn|>' -}}
|
| 21 |
+
{%- endfor -%}
|
| 22 |
+
{%- if add_generation_prompt -%}
|
| 23 |
+
{{- '<|think|>' if enable_thinking is defined and enable_thinking else '<|assistant|>' -}}
|
| 24 |
+
{%- endif -%}
|
config.json
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"apply_legacy_rope_quirk": false,
|
| 3 |
+
"architectures": [
|
| 4 |
+
"AutoModelForCausalLM"
|
| 5 |
+
],
|
| 6 |
+
"attention_mechanism": null,
|
| 7 |
+
"attn_pdrop": 0.0,
|
| 8 |
+
"auto_map": {
|
| 9 |
+
"AutoConfig": "han2han_config.Han2HanConfig",
|
| 10 |
+
"AutoModel": "modeling_han2han.Han2Han",
|
| 11 |
+
"AutoModelForMultipleChoice": "modeling_han2han.Han2HanForMultipleChoice",
|
| 12 |
+
"AutoModelForQuestionAnswering": "modeling_han2han.Han2HanForQuestionAnswering",
|
| 13 |
+
"AutoModelForSeq2SeqLM": "modeling_han2han.Han2Han",
|
| 14 |
+
"AutoModelForSequenceClassification": "modeling_han2han.Han2HanForSequenceClassification",
|
| 15 |
+
"AutoModelForTokenClassification": "modeling_han2han.Han2HanForTokenClassification",
|
| 16 |
+
"AutoModelForCausalLM": "modeling_han2han.Han2HanForCausalLM"
|
| 17 |
+
},
|
| 18 |
+
"bos_token_id": 2,
|
| 19 |
+
"char_is_unified_cjk": false,
|
| 20 |
+
"char_subwords": true,
|
| 21 |
+
"char_vocab_size": 5376,
|
| 22 |
+
"classf_pdrop": 0.1,
|
| 23 |
+
"classifier_head_type": "linear",
|
| 24 |
+
"cross_attn_num_heads": 4,
|
| 25 |
+
"cross_attn_num_kv_heads": 2,
|
| 26 |
+
"cross_attn_pdrop": 0.0,
|
| 27 |
+
"d_ff": 2048,
|
| 28 |
+
"d_model": 640,
|
| 29 |
+
"d_prime": 1024,
|
| 30 |
+
"decoder_attention_types": [
|
| 31 |
+
"mha-sliding",
|
| 32 |
+
"mha-sliding",
|
| 33 |
+
"mha-sliding",
|
| 34 |
+
"mha-sliding",
|
| 35 |
+
"mha-sliding",
|
| 36 |
+
"mha"
|
| 37 |
+
],
|
| 38 |
+
"decoder_cross_attention_types": [
|
| 39 |
+
"mha"
|
| 40 |
+
],
|
| 41 |
+
"decoder_nlayer": 18,
|
| 42 |
+
"decoder_norm_type": "rmsnorm",
|
| 43 |
+
"decoder_start_token_id": 3,
|
| 44 |
+
"dense_ffn_activation": "geglu",
|
| 45 |
+
"dtype": "float32",
|
| 46 |
+
"embd_pdrop": 0.0,
|
| 47 |
+
"embedding_dropout_rate": 0.0,
|
| 48 |
+
"encoder_attention_types": [
|
| 49 |
+
"mha-sliding",
|
| 50 |
+
"mha-sliding",
|
| 51 |
+
"mha-sliding",
|
| 52 |
+
"mha-sliding",
|
| 53 |
+
"mha-sliding",
|
| 54 |
+
"mha"
|
| 55 |
+
],
|
| 56 |
+
"encoder_nlayer": 18,
|
| 57 |
+
"encoder_norm_type": "rmsnorm",
|
| 58 |
+
"eos_token_id": 3,
|
| 59 |
+
"ffn_activation": "geglu",
|
| 60 |
+
"head_dim": 256,
|
| 61 |
+
"init_biases_normal": true,
|
| 62 |
+
"initializer_range": 0.02,
|
| 63 |
+
"is_encoder_decoder": true,
|
| 64 |
+
"jamo_subwords": true,
|
| 65 |
+
"jamo_vocab_size": 4992,
|
| 66 |
+
"kernel_init_scale": 0.1,
|
| 67 |
+
"kernel_init_type": "variance_scaling",
|
| 68 |
+
"label_smoothing": 0.0,
|
| 69 |
+
"layer_norm_epsilon": 1e-05,
|
| 70 |
+
"layer_pdrop": 0.0,
|
| 71 |
+
"model_type": "han2han",
|
| 72 |
+
"n_positions": 2048,
|
| 73 |
+
"num_heads": 4,
|
| 74 |
+
"num_kv_heads": 1,
|
| 75 |
+
"pad_token_id": 0,
|
| 76 |
+
"qk_norm_post_rope": true,
|
| 77 |
+
"query_pre_attn_scalar": 256,
|
| 78 |
+
"remat_policy": "full",
|
| 79 |
+
"resid_pdrop": 0.0,
|
| 80 |
+
"return_dict": true,
|
| 81 |
+
"rope_theta": 500000,
|
| 82 |
+
"rope_theta_sliding": 10000,
|
| 83 |
+
"seed": 42,
|
| 84 |
+
"sft_decoder_start_token_id_default": null,
|
| 85 |
+
"sft_decoder_start_token_id_thinking": null,
|
| 86 |
+
"sft_eos_token_id": null,
|
| 87 |
+
"sliding_window_size": 128,
|
| 88 |
+
"subword_embed_dim": 384,
|
| 89 |
+
"subword_entry_fusion": true,
|
| 90 |
+
"entry_pool_mode": "legacy",
|
| 91 |
+
"entry_pad_bias_slots": 0,
|
| 92 |
+
"entry_slot_width": 32,
|
| 93 |
+
"entry_pool_unroll": true,
|
| 94 |
+
"tie_encoder_decoder": true,
|
| 95 |
+
"tie_input_output_embeddings": true,
|
| 96 |
+
"tie_subtoken_embeddings": false,
|
| 97 |
+
"tie_word_embeddings": true,
|
| 98 |
+
"transformers_version": "5.18.0",
|
| 99 |
+
"use_bart_collator": true,
|
| 100 |
+
"use_bart_training": true,
|
| 101 |
+
"use_bias": true,
|
| 102 |
+
"use_fla_fused_mlp": false,
|
| 103 |
+
"use_fla_fused_norm": false,
|
| 104 |
+
"use_fla_fused_rotary": false,
|
| 105 |
+
"use_han2han_transcription": "false",
|
| 106 |
+
"use_learned_bidirectional": true,
|
| 107 |
+
"use_qk_norm": true,
|
| 108 |
+
"use_scan_layers": true,
|
| 109 |
+
"use_sub_ln": true,
|
| 110 |
+
"vocab_size": 38400
|
| 111 |
+
}
|
generation_config.json
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token_id": 2,
|
| 3 |
+
"decoder_start_token_id": 9,
|
| 4 |
+
"eos_token_id": 10,
|
| 5 |
+
"pad_token_id": 0,
|
| 6 |
+
"use_cache": true,
|
| 7 |
+
"transformers_version": "5.18.0"
|
| 8 |
+
}
|
han2han_config.py
ADDED
|
@@ -0,0 +1,565 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
# coding: utf-8
|
| 3 |
+
|
| 4 |
+
import logging
|
| 5 |
+
|
| 6 |
+
from transformers.configuration_utils import PretrainedConfig
|
| 7 |
+
from typing import Optional, List
|
| 8 |
+
|
| 9 |
+
logger = logging.getLogger(__name__)
|
| 10 |
+
|
| 11 |
+
class Han2HanConfig(PretrainedConfig):
|
| 12 |
+
model_type = "han2han"
|
| 13 |
+
keys_to_ignore_at_inference = ["past_key_values"]
|
| 14 |
+
has_no_defaults_at_init = True
|
| 15 |
+
attribute_map = {
|
| 16 |
+
"hidden_size": "d_model",
|
| 17 |
+
"num_attention_heads": "num_heads",
|
| 18 |
+
"num_hidden_layers": "decoder_nlayer", # HF expects this for decoder
|
| 19 |
+
"num_layers": "decoder_nlayer",
|
| 20 |
+
"encoder_layers": "encoder_nlayer",
|
| 21 |
+
"decoder_layers": "decoder_nlayer",
|
| 22 |
+
"d_kv": "d_prime",
|
| 23 |
+
"intermediate_size": "d_ff",
|
| 24 |
+
"hidden_dropout_prob": "resid_pdrop",
|
| 25 |
+
"attention_probs_dropout_prob": "attn_pdrop",
|
| 26 |
+
}
|
| 27 |
+
|
| 28 |
+
# Fields that should always be saved, even if they match PretrainedConfig defaults
|
| 29 |
+
_always_save = ["tie_word_embeddings", "tie_encoder_decoder"]
|
| 30 |
+
|
| 31 |
+
# Explicitly list attributes that should always be saved
|
| 32 |
+
def to_diff_dict(self):
|
| 33 |
+
"""
|
| 34 |
+
Override to ensure certain fields are always saved, even if they match
|
| 35 |
+
PretrainedConfig defaults.
|
| 36 |
+
"""
|
| 37 |
+
# Get the diff dict from parent
|
| 38 |
+
output = super().to_diff_dict()
|
| 39 |
+
|
| 40 |
+
# Force inclusion of fields we always want to save
|
| 41 |
+
for field in self._always_save:
|
| 42 |
+
if hasattr(self, field):
|
| 43 |
+
output[field] = getattr(self, field)
|
| 44 |
+
|
| 45 |
+
return output
|
| 46 |
+
|
| 47 |
+
def get_text_config(self, decoder=None, encoder=None):
|
| 48 |
+
"""Return this config for either side of the model.
|
| 49 |
+
|
| 50 |
+
The encoder and decoder share one flat config. The base implementation
|
| 51 |
+
treats flat encoder-decoder configs as legacy and strips the `decoder_`
|
| 52 |
+
prefix from every key (`decoder_nlayer` -> `nlayer`), which leaves
|
| 53 |
+
`num_hidden_layers` unresolvable when `generate()` sizes its KV cache.
|
| 54 |
+
"""
|
| 55 |
+
return self
|
| 56 |
+
|
| 57 |
+
def __init__(
|
| 58 |
+
self,
|
| 59 |
+
|
| 60 |
+
jamo_subwords: bool = False,
|
| 61 |
+
char_subwords: bool = False,
|
| 62 |
+
use_han2han_transcription: str = 'false', # 'false', 'true', 'hangul_only', 'reverse'
|
| 63 |
+
|
| 64 |
+
vocab_size: int = 38000,
|
| 65 |
+
jamo_vocab_size: int = 8000,
|
| 66 |
+
char_vocab_size: int = 8000,
|
| 67 |
+
subword_embed_dim: Optional[int] = None, # hidden dim of wje/wce subtoken embeddings; None defaults to d_model // 2
|
| 68 |
+
char_is_unified_cjk: bool = False,
|
| 69 |
+
decoder_nlayer: int = 4,
|
| 70 |
+
encoder_nlayer: int = 4,
|
| 71 |
+
n_positions: int = 1024,
|
| 72 |
+
d_model: int = 768,
|
| 73 |
+
d_prime: Optional[int] = None,
|
| 74 |
+
rope_theta: float = 10000.0,
|
| 75 |
+
rope_theta_sliding: Optional[float] = None, # per-layer theta override for 'sliding'/'local' attention layers (None = inherit rope_theta)
|
| 76 |
+
apply_legacy_rope_quirk: Optional[bool] = None, # None = auto-detect legacy scan_rope_theta quirk; False = opt out (use rope_theta/rope_theta_sliding as specified, for V2+); True = force-apply
|
| 77 |
+
d_ff: int = 3072,
|
| 78 |
+
attention_mechanism: str = 'mha', # global default attention type for all layers
|
| 79 |
+
encoder_attention_types: Optional[List[str]] = None, # per-layer encoder self-attention types
|
| 80 |
+
decoder_attention_types: Optional[List[str]] = None, # per-layer decoder self-attention types
|
| 81 |
+
decoder_cross_attention_types: Optional[List[str]] = None, # per-layer decoder cross-attention types
|
| 82 |
+
sliding_window_size: int = 256, # window size for 'mha-sliding' layers (0 = full attention)
|
| 83 |
+
ffn_activation: str = "swiglu", # 'swiglu', 'geglu', 'reglu2', 'gelu', 'gelu_new', 'relu2'
|
| 84 |
+
dense_ffn_activation: Optional[str] = None, # activation override; None = follow ffn_activation
|
| 85 |
+
use_fla_fused_mlp: bool = False,
|
| 86 |
+
use_fla_fused_norm: bool = False,
|
| 87 |
+
use_fla_fused_rotary: bool = False,
|
| 88 |
+
num_heads: Optional[int] = None,
|
| 89 |
+
head_dim: Optional[int] = None, # MHA only: per-head dim. If both set with d_prime, must be consistent. Required for GQA.
|
| 90 |
+
num_kv_heads: Optional[int] = None, # MHA only: KV heads for self-attn (None = num_heads = full MHA). 1 = MQA. Must divide num_heads.
|
| 91 |
+
cross_attn_num_heads: Optional[int] = None, # MHA only: Q heads for cross-attn (None = num_heads). Lets cross-attn use full d_model fidelity (cross_attn_num_heads * head_dim) while self-attn stays compressed.
|
| 92 |
+
cross_attn_num_kv_heads: Optional[int] = None, # MHA only: KV heads for cross-attn (None = num_kv_heads). Must divide cross_attn_num_heads.
|
| 93 |
+
use_qk_norm: bool = False, # MHA only: per-head RMSNorm on Q and K (Gemma 3 / T5Gemma 2 style)
|
| 94 |
+
query_pre_attn_scalar: Optional[float] = None, # MHA only: HF-Gemma 3 semantics. Q multiplier = scalar ** -0.5. None = head_dim ** -0.5.
|
| 95 |
+
num_labels: int = 3,
|
| 96 |
+
use_learned_bidirectional: float = True,
|
| 97 |
+
remat_policy: str = "full",
|
| 98 |
+
layer_pdrop: float = 0.1,
|
| 99 |
+
resid_pdrop: float = 0.1,
|
| 100 |
+
embd_pdrop: float = 0.1,
|
| 101 |
+
attn_pdrop: float = 0.1,
|
| 102 |
+
cross_attn_pdrop: float = 0.15,
|
| 103 |
+
classf_pdrop: float = 0.1,
|
| 104 |
+
classifier_head_type: str = 'linear', # 'linear' (T5Gemma 2) or 'mlp' (RoBERTa-style tanh)
|
| 105 |
+
embedding_dropout_rate: float = 0.0, # Probability of dropping each type of embedding module (wte/wje/wce)
|
| 106 |
+
layer_norm_epsilon: float = 1e-5,
|
| 107 |
+
decoder_norm_type: str = 'rmsnorm', # 'rmsnorm', 'rmsnorm_bias', 'layernorm'
|
| 108 |
+
encoder_norm_type: str = 'rmsnorm', # 'rmsnorm', 'rmsnorm_bias', 'layernorm'
|
| 109 |
+
initializer_range: float = 0.02,
|
| 110 |
+
kernel_init_type: str = 'normal', # 'normal' or 'variance_scaling'
|
| 111 |
+
kernel_init_scale: float = 0.1, # scale for variance_scaling (0.1=Switch-style, 1.0=lecun)
|
| 112 |
+
init_biases_normal: bool = False, # if True, init biases as normal(stddev=initializer_range) (V1 behavior); else zeros
|
| 113 |
+
init_cache: bool = False,
|
| 114 |
+
pad_token_id: int = 1,
|
| 115 |
+
decoder_start_token_id: int = 0,
|
| 116 |
+
tie_word_embeddings: bool = True,
|
| 117 |
+
tie_encoder_decoder: bool = False,
|
| 118 |
+
tie_input_output_embeddings: bool = False,
|
| 119 |
+
tie_subtoken_embeddings: bool = True,
|
| 120 |
+
return_dict: bool = True,
|
| 121 |
+
seed: int = 0,
|
| 122 |
+
eos_token_id: int = 2,
|
| 123 |
+
bos_token_id: int = 0,
|
| 124 |
+
use_bart_training: bool = True,
|
| 125 |
+
use_bart_collator: bool = True,
|
| 126 |
+
|
| 127 |
+
# SFT-only token ids. Read by ChatSFTCollator; not used by pretraining.
|
| 128 |
+
# Decoder is primed with <|think|> when reasoning is enabled, otherwise
|
| 129 |
+
# with <|assistant|>; turn closes with <|end_of_turn|>. None means look
|
| 130 |
+
# up from tokenizer at collator init time (recommended).
|
| 131 |
+
sft_decoder_start_token_id_thinking: Optional[int] = None,
|
| 132 |
+
sft_decoder_start_token_id_default: Optional[int] = None,
|
| 133 |
+
sft_eos_token_id: Optional[int] = None,
|
| 134 |
+
|
| 135 |
+
use_sub_ln: bool = False, # SubLN: RMSNorm before output projections in attn and FFN
|
| 136 |
+
|
| 137 |
+
# Bias configuration
|
| 138 |
+
use_bias: bool = True, # global toggle for biases in all linear layers
|
| 139 |
+
|
| 140 |
+
label_smoothing: float = 0.1, # label smoothing alpha for cross-entropy loss
|
| 141 |
+
|
| 142 |
+
# scan layers (compile one layer body, repeat N times via jax.lax.scan)
|
| 143 |
+
use_scan_layers: bool = False,
|
| 144 |
+
|
| 145 |
+
**kwargs,
|
| 146 |
+
):
|
| 147 |
+
self.vocab_size = vocab_size
|
| 148 |
+
self.jamo_vocab_size = jamo_vocab_size
|
| 149 |
+
self.char_vocab_size = char_vocab_size
|
| 150 |
+
# subtoken embedding hidden dim: defaults to d_model // 2 when unspecified
|
| 151 |
+
self.subword_embed_dim = subword_embed_dim if subword_embed_dim is not None else d_model // 2
|
| 152 |
+
self.char_is_unified_cjk = char_is_unified_cjk
|
| 153 |
+
self.jamo_subwords = jamo_subwords
|
| 154 |
+
self.char_subwords = char_subwords
|
| 155 |
+
self.use_han2han_transcription = use_han2han_transcription
|
| 156 |
+
self.decoder_nlayer = decoder_nlayer
|
| 157 |
+
self.encoder_nlayer = encoder_nlayer
|
| 158 |
+
self.n_positions = n_positions
|
| 159 |
+
self.d_model = d_model
|
| 160 |
+
self.rope_theta = rope_theta
|
| 161 |
+
self.rope_theta_sliding = rope_theta_sliding
|
| 162 |
+
self.apply_legacy_rope_quirk = apply_legacy_rope_quirk
|
| 163 |
+
self.d_prime = d_prime
|
| 164 |
+
self.d_ff = d_ff
|
| 165 |
+
self.ffn_activation = ffn_activation
|
| 166 |
+
self.dense_ffn_activation = dense_ffn_activation if dense_ffn_activation is not None else ffn_activation
|
| 167 |
+
# Only set num_heads for MHA attention
|
| 168 |
+
self.num_heads = num_heads if attention_mechanism == 'mha' or any(
|
| 169 |
+
'mha' in attn_type for attn_type in (encoder_attention_types or [])) else None
|
| 170 |
+
self.head_dim = head_dim
|
| 171 |
+
self.num_kv_heads = num_kv_heads
|
| 172 |
+
self.cross_attn_num_heads = cross_attn_num_heads
|
| 173 |
+
self.cross_attn_num_kv_heads = cross_attn_num_kv_heads
|
| 174 |
+
self.use_qk_norm = use_qk_norm
|
| 175 |
+
self.query_pre_attn_scalar = query_pre_attn_scalar
|
| 176 |
+
self.use_fla_fused_mlp = use_fla_fused_mlp
|
| 177 |
+
self.use_fla_fused_norm = use_fla_fused_norm
|
| 178 |
+
self.use_fla_fused_rotary = use_fla_fused_rotary
|
| 179 |
+
self.remat_policy = remat_policy
|
| 180 |
+
self.layer_pdrop = layer_pdrop
|
| 181 |
+
self.resid_pdrop = resid_pdrop
|
| 182 |
+
self.embd_pdrop = embd_pdrop
|
| 183 |
+
self.attn_pdrop = attn_pdrop
|
| 184 |
+
self.cross_attn_pdrop = cross_attn_pdrop
|
| 185 |
+
self.classf_pdrop = classf_pdrop
|
| 186 |
+
self.classifier_head_type = classifier_head_type
|
| 187 |
+
self.embedding_dropout_rate = embedding_dropout_rate
|
| 188 |
+
self.layer_norm_epsilon = layer_norm_epsilon
|
| 189 |
+
self.decoder_norm_type = decoder_norm_type
|
| 190 |
+
self.encoder_norm_type = encoder_norm_type
|
| 191 |
+
self.initializer_range = initializer_range
|
| 192 |
+
self.kernel_init_type = kernel_init_type
|
| 193 |
+
self.kernel_init_scale = kernel_init_scale
|
| 194 |
+
self.init_biases_normal = init_biases_normal
|
| 195 |
+
self.use_cache = init_cache
|
| 196 |
+
self.tie_input_output_embeddings = tie_input_output_embeddings
|
| 197 |
+
self.use_sub_ln = use_sub_ln
|
| 198 |
+
|
| 199 |
+
# bias configuration
|
| 200 |
+
self.use_bias = use_bias
|
| 201 |
+
|
| 202 |
+
self.label_smoothing = label_smoothing
|
| 203 |
+
self.use_scan_layers = use_scan_layers
|
| 204 |
+
|
| 205 |
+
super().__init__(
|
| 206 |
+
pad_token_id = pad_token_id,
|
| 207 |
+
decoder_start_token_id = decoder_start_token_id,
|
| 208 |
+
tie_encoder_decoder = tie_encoder_decoder,
|
| 209 |
+
tie_word_embeddings = tie_word_embeddings,
|
| 210 |
+
bos_token_id = bos_token_id,
|
| 211 |
+
eos_token_id = eos_token_id,
|
| 212 |
+
**kwargs
|
| 213 |
+
)
|
| 214 |
+
|
| 215 |
+
self.num_labels = num_labels
|
| 216 |
+
self.is_encoder_decoder = True
|
| 217 |
+
self.tie_word_embeddings = tie_word_embeddings
|
| 218 |
+
self.tie_encoder_decoder = tie_encoder_decoder
|
| 219 |
+
self.tie_subtoken_embeddings = tie_subtoken_embeddings
|
| 220 |
+
self.seed = seed
|
| 221 |
+
self.return_dict = return_dict
|
| 222 |
+
self.eos_token_id = eos_token_id
|
| 223 |
+
self.sft_decoder_start_token_id_thinking = sft_decoder_start_token_id_thinking
|
| 224 |
+
self.sft_decoder_start_token_id_default = sft_decoder_start_token_id_default
|
| 225 |
+
self.sft_eos_token_id = sft_eos_token_id
|
| 226 |
+
self.attention_mechanism = attention_mechanism
|
| 227 |
+
self.encoder_attention_types = encoder_attention_types
|
| 228 |
+
self.decoder_attention_types = decoder_attention_types
|
| 229 |
+
self.decoder_cross_attention_types = decoder_cross_attention_types
|
| 230 |
+
self.sliding_window_size = sliding_window_size
|
| 231 |
+
self.use_learned_bidirectional = use_learned_bidirectional
|
| 232 |
+
self.use_bart_training = use_bart_training
|
| 233 |
+
self.use_bart_collator = use_bart_collator
|
| 234 |
+
|
| 235 |
+
self._sanitize_d_prime(encoder_attention_types, decoder_attention_types, decoder_cross_attention_types)
|
| 236 |
+
self._resolve_mha_attention_fields(encoder_attention_types, decoder_attention_types, decoder_cross_attention_types)
|
| 237 |
+
|
| 238 |
+
if self.d_prime > 256 and self.use_fla_fused_rotary:
|
| 239 |
+
logger.warning(
|
| 240 |
+
"WARNING: `d_prime` is greater than 256 and `use_fla_fused_rotary` is set to `True`. "
|
| 241 |
+
"This is not supported. Falling back to non-fused rotary embeddings."
|
| 242 |
+
)
|
| 243 |
+
self.use_fla_fused_rotary = False
|
| 244 |
+
|
| 245 |
+
self._validate_attention_config()
|
| 246 |
+
|
| 247 |
+
if self.use_scan_layers:
|
| 248 |
+
self._validate_scan_config()
|
| 249 |
+
|
| 250 |
+
self._apply_legacy_rope_quirk()
|
| 251 |
+
|
| 252 |
+
def _apply_legacy_rope_quirk(self) -> None:
|
| 253 |
+
"""Rewrite rope_theta when the legacy Flax `_scan_key` quirk applied.
|
| 254 |
+
|
| 255 |
+
Pre-fix Flax `_identify_scan_groups` collapsed all MHA variants
|
| 256 |
+
('mha', 'mha-sliding', 'mha-local') under a single scan key. So when
|
| 257 |
+
`use_scan_layers=True` and the attention pattern mixes 'mha' with
|
| 258 |
+
'mha-sliding' or 'mha-local', all layers in the resulting scan stack
|
| 259 |
+
share one `FlaxHan2HanBlock` instance whose `effective_rope_theta` is
|
| 260 |
+
baked from `position_specs[0]`. In practice the configured
|
| 261 |
+
`rope_theta` field was overwritten by `rope_theta_sliding` whenever
|
| 262 |
+
the first attention layer was sliding.
|
| 263 |
+
|
| 264 |
+
Checkpoints trained under that quirk have a saved `rope_theta` that
|
| 265 |
+
doesn't reflect what the weights actually saw. This sanitizer detects
|
| 266 |
+
the pattern and rewrites the in-memory config so the runtime loads
|
| 267 |
+
the model with values matching its training behavior. Flax post-fix
|
| 268 |
+
resolves `effective_rope_theta` per call from `override_window_size`,
|
| 269 |
+
so the sanitizer's rewrite is what makes legacy checkpoints behave
|
| 270 |
+
identically under the corrected dispatch path.
|
| 271 |
+
|
| 272 |
+
Honors `self.apply_legacy_rope_quirk`:
|
| 273 |
+
- None (default): auto-apply when the heuristic triggers (existing
|
| 274 |
+
checkpoints, no opt-out specified -- safe default).
|
| 275 |
+
- False: explicit opt-out (V2+ runs that want hybrid rope_theta).
|
| 276 |
+
Sanitizer becomes a no-op even when the trigger pattern matches.
|
| 277 |
+
- True: force-apply (mostly for testing).
|
| 278 |
+
"""
|
| 279 |
+
opt = getattr(self, 'apply_legacy_rope_quirk', None)
|
| 280 |
+
if opt is False:
|
| 281 |
+
return
|
| 282 |
+
|
| 283 |
+
if not getattr(self, 'use_scan_layers', False):
|
| 284 |
+
return
|
| 285 |
+
if getattr(self, 'rope_theta_sliding', None) is None:
|
| 286 |
+
return
|
| 287 |
+
if self.rope_theta == self.rope_theta_sliding:
|
| 288 |
+
return
|
| 289 |
+
|
| 290 |
+
def _has_mixed_mha(types):
|
| 291 |
+
if not types:
|
| 292 |
+
return False
|
| 293 |
+
has_sliding = any(
|
| 294 |
+
('sliding' in t) or ('local' in t)
|
| 295 |
+
for t in types if t
|
| 296 |
+
)
|
| 297 |
+
has_full_mha = any(t == 'mha' for t in types)
|
| 298 |
+
return has_sliding and has_full_mha
|
| 299 |
+
|
| 300 |
+
triggered = opt is True or (
|
| 301 |
+
_has_mixed_mha(self.encoder_attention_types)
|
| 302 |
+
or _has_mixed_mha(self.decoder_attention_types)
|
| 303 |
+
)
|
| 304 |
+
if not triggered:
|
| 305 |
+
return
|
| 306 |
+
|
| 307 |
+
logger.warning(
|
| 308 |
+
"[Han2HanConfig] Detected legacy scan_rope_theta quirk "
|
| 309 |
+
"(use_scan_layers=True + mixed mha/mha-sliding pattern + "
|
| 310 |
+
"rope_theta != rope_theta_sliding). Pre-fix Flax "
|
| 311 |
+
"_identify_scan_groups collapsed MHA variants into one scan "
|
| 312 |
+
"stack whose rope_theta was baked from position_specs[0], so "
|
| 313 |
+
"the saved rope_theta=%s was overwritten by rope_theta_sliding=%s "
|
| 314 |
+
"during training. Rewriting in-memory rope_theta=%s to match "
|
| 315 |
+
"actual training behavior. To opt out for new training, set "
|
| 316 |
+
"apply_legacy_rope_quirk=False (CLI: --no_apply_legacy_rope_quirk).",
|
| 317 |
+
self.rope_theta, self.rope_theta_sliding, self.rope_theta_sliding,
|
| 318 |
+
)
|
| 319 |
+
self.rope_theta = self.rope_theta_sliding
|
| 320 |
+
|
| 321 |
+
def _sanitize_d_prime(
|
| 322 |
+
self,
|
| 323 |
+
encoder_attention_types: Optional[List[str]],
|
| 324 |
+
decoder_attention_types: Optional[List[str]],
|
| 325 |
+
decoder_cross_attention_types: Optional[List[str]],
|
| 326 |
+
) -> None:
|
| 327 |
+
"""Sanitize d_prime based on attention configuration.
|
| 328 |
+
|
| 329 |
+
- Layerwise config: d_prime defaults to d_model if not specified
|
| 330 |
+
- Global MHA: d_prime defaults to d_model if not specified
|
| 331 |
+
"""
|
| 332 |
+
has_layerwise = (
|
| 333 |
+
encoder_attention_types is not None or
|
| 334 |
+
decoder_attention_types is not None or
|
| 335 |
+
decoder_cross_attention_types is not None
|
| 336 |
+
)
|
| 337 |
+
|
| 338 |
+
if has_layerwise:
|
| 339 |
+
if self.d_prime is None and self.head_dim is None:
|
| 340 |
+
self.d_prime = self.d_model
|
| 341 |
+
logger.info(f"d_prime not specified for layerwise config, defaulting to d_model={self.d_model}")
|
| 342 |
+
elif self.attention_mechanism == 'mha':
|
| 343 |
+
if self.d_prime is None and self.head_dim is None:
|
| 344 |
+
self.d_prime = self.d_model
|
| 345 |
+
logger.info(f"d_prime not specified for MHA, defaulting to d_model={self.d_model}")
|
| 346 |
+
else:
|
| 347 |
+
if self.d_prime is None:
|
| 348 |
+
raise ValueError(
|
| 349 |
+
f"d_prime must be explicitly set for attention_mechanism='{self.attention_mechanism}'."
|
| 350 |
+
)
|
| 351 |
+
|
| 352 |
+
def _resolve_mha_attention_fields(
|
| 353 |
+
self,
|
| 354 |
+
encoder_attention_types: Optional[List[str]],
|
| 355 |
+
decoder_attention_types: Optional[List[str]],
|
| 356 |
+
decoder_cross_attention_types: Optional[List[str]],
|
| 357 |
+
) -> None:
|
| 358 |
+
"""Resolve head_dim/d_prime consistency and num_kv_heads defaults for MHA."""
|
| 359 |
+
all_types = []
|
| 360 |
+
if encoder_attention_types is not None:
|
| 361 |
+
all_types.extend(encoder_attention_types)
|
| 362 |
+
if decoder_attention_types is not None:
|
| 363 |
+
all_types.extend(decoder_attention_types)
|
| 364 |
+
if decoder_cross_attention_types is not None:
|
| 365 |
+
all_types.extend(t for t in decoder_cross_attention_types if t is not None)
|
| 366 |
+
has_layerwise = bool(all_types)
|
| 367 |
+
has_mha = (
|
| 368 |
+
any(t == 'mha' or t.startswith('mha-') for t in all_types)
|
| 369 |
+
if has_layerwise
|
| 370 |
+
else self.attention_mechanism == 'mha'
|
| 371 |
+
)
|
| 372 |
+
|
| 373 |
+
if not has_mha:
|
| 374 |
+
return
|
| 375 |
+
|
| 376 |
+
if self.num_heads is None:
|
| 377 |
+
raise ValueError(
|
| 378 |
+
"num_heads must be specified for MHA attention."
|
| 379 |
+
)
|
| 380 |
+
|
| 381 |
+
if self.head_dim is not None and self.d_prime is not None:
|
| 382 |
+
if self.head_dim * self.num_heads != self.d_prime:
|
| 383 |
+
raise ValueError(
|
| 384 |
+
f"head_dim ({self.head_dim}) * num_heads ({self.num_heads}) = "
|
| 385 |
+
f"{self.head_dim * self.num_heads} does not match d_prime ({self.d_prime})."
|
| 386 |
+
)
|
| 387 |
+
elif self.head_dim is not None:
|
| 388 |
+
self.d_prime = self.head_dim * self.num_heads
|
| 389 |
+
elif self.d_prime is not None:
|
| 390 |
+
if self.d_prime % self.num_heads != 0:
|
| 391 |
+
raise ValueError(
|
| 392 |
+
f"d_prime ({self.d_prime}) must be divisible by num_heads ({self.num_heads})."
|
| 393 |
+
)
|
| 394 |
+
self.head_dim = self.d_prime // self.num_heads
|
| 395 |
+
else:
|
| 396 |
+
raise ValueError(
|
| 397 |
+
"MHA requires at least one of head_dim or d_prime to be set."
|
| 398 |
+
)
|
| 399 |
+
|
| 400 |
+
if self.num_kv_heads is None:
|
| 401 |
+
self.num_kv_heads = self.num_heads
|
| 402 |
+
if self.num_kv_heads < 1 or self.num_heads % self.num_kv_heads != 0:
|
| 403 |
+
raise ValueError(
|
| 404 |
+
f"num_kv_heads ({self.num_kv_heads}) must be >= 1 and divide "
|
| 405 |
+
f"num_heads ({self.num_heads}) evenly."
|
| 406 |
+
)
|
| 407 |
+
|
| 408 |
+
if self.cross_attn_num_heads is None:
|
| 409 |
+
self.cross_attn_num_heads = self.num_heads
|
| 410 |
+
if self.cross_attn_num_heads < 1:
|
| 411 |
+
raise ValueError(
|
| 412 |
+
f"cross_attn_num_heads ({self.cross_attn_num_heads}) must be >= 1."
|
| 413 |
+
)
|
| 414 |
+
|
| 415 |
+
if self.cross_attn_num_kv_heads is None:
|
| 416 |
+
self.cross_attn_num_kv_heads = self.num_kv_heads
|
| 417 |
+
if self.cross_attn_num_kv_heads < 1 or self.cross_attn_num_heads % self.cross_attn_num_kv_heads != 0:
|
| 418 |
+
raise ValueError(
|
| 419 |
+
f"cross_attn_num_kv_heads ({self.cross_attn_num_kv_heads}) must be >= 1 and divide "
|
| 420 |
+
f"cross_attn_num_heads ({self.cross_attn_num_heads}) evenly."
|
| 421 |
+
)
|
| 422 |
+
|
| 423 |
+
if self.query_pre_attn_scalar is not None and self.query_pre_attn_scalar <= 0:
|
| 424 |
+
raise ValueError(
|
| 425 |
+
f"query_pre_attn_scalar must be positive, got {self.query_pre_attn_scalar}."
|
| 426 |
+
)
|
| 427 |
+
|
| 428 |
+
def _validate_attention_config(self) -> None:
|
| 429 |
+
"""Validate mutual exclusivity of global vs layerwise attention config."""
|
| 430 |
+
has_layerwise = (
|
| 431 |
+
self.encoder_attention_types is not None or
|
| 432 |
+
self.decoder_attention_types is not None or
|
| 433 |
+
self.decoder_cross_attention_types is not None
|
| 434 |
+
)
|
| 435 |
+
|
| 436 |
+
if has_layerwise:
|
| 437 |
+
missing = []
|
| 438 |
+
if self.encoder_attention_types is None:
|
| 439 |
+
missing.append('encoder_attention_types')
|
| 440 |
+
if self.decoder_attention_types is None:
|
| 441 |
+
missing.append('decoder_attention_types')
|
| 442 |
+
if self.decoder_cross_attention_types is None:
|
| 443 |
+
missing.append('decoder_cross_attention_types')
|
| 444 |
+
|
| 445 |
+
if missing:
|
| 446 |
+
raise ValueError(
|
| 447 |
+
f"Per-layer attention mode requires all three lists to be specified. "
|
| 448 |
+
f"Missing: {missing}. Either specify all layerwise types, or remove "
|
| 449 |
+
f"all and use 'attention_mechanism' for a global default."
|
| 450 |
+
)
|
| 451 |
+
|
| 452 |
+
if self.attention_mechanism is not None:
|
| 453 |
+
logger.warning(
|
| 454 |
+
f"Both attention_mechanism ('{self.attention_mechanism}') and layerwise "
|
| 455 |
+
f"attention types are set. Layerwise types take precedence; "
|
| 456 |
+
f"attention_mechanism will be ignored."
|
| 457 |
+
)
|
| 458 |
+
else:
|
| 459 |
+
if self.attention_mechanism is None:
|
| 460 |
+
raise ValueError(
|
| 461 |
+
"No attention configuration specified. Either set 'attention_mechanism' "
|
| 462 |
+
"for a global default, or specify all three layerwise types: "
|
| 463 |
+
"encoder_attention_types, decoder_attention_types, decoder_cross_attention_types."
|
| 464 |
+
)
|
| 465 |
+
|
| 466 |
+
def _validate_scan_config(self) -> None:
|
| 467 |
+
"""Validate config constraints for scanned layer execution."""
|
| 468 |
+
# layerdrop breaks scan (stochastic layer skip)
|
| 469 |
+
if self.layer_pdrop > 0:
|
| 470 |
+
raise ValueError(
|
| 471 |
+
f"use_scan_layers is incompatible with layerdrop > 0. "
|
| 472 |
+
f"Got layer_pdrop={self.layer_pdrop}. Set layerdrop: 0.0."
|
| 473 |
+
)
|
| 474 |
+
|
| 475 |
+
def make_kernel_init(self, dtype=None):
|
| 476 |
+
"""Create kernel initializer based on config."""
|
| 477 |
+
# try block keeps transformers' remote-code import check from requiring flax
|
| 478 |
+
try:
|
| 479 |
+
from flax import nnx
|
| 480 |
+
except ImportError:
|
| 481 |
+
raise
|
| 482 |
+
kwargs = {'dtype': dtype} if dtype is not None else {}
|
| 483 |
+
if self.kernel_init_type == 'variance_scaling':
|
| 484 |
+
return nnx.initializers.variance_scaling(
|
| 485 |
+
self.kernel_init_scale, 'fan_in', 'truncated_normal', **kwargs
|
| 486 |
+
)
|
| 487 |
+
return nnx.initializers.normal(stddev=self.initializer_range, **kwargs)
|
| 488 |
+
|
| 489 |
+
def make_bias_init(self):
|
| 490 |
+
"""Create bias initializer based on config.
|
| 491 |
+
|
| 492 |
+
V1 behavior is normal(stddev=initializer_range); current default is zeros.
|
| 493 |
+
"""
|
| 494 |
+
# try block keeps transformers' remote-code import check from requiring flax
|
| 495 |
+
try:
|
| 496 |
+
from flax import nnx
|
| 497 |
+
except ImportError:
|
| 498 |
+
raise
|
| 499 |
+
if self.init_biases_normal:
|
| 500 |
+
return nnx.initializers.normal(stddev=self.initializer_range)
|
| 501 |
+
return nnx.initializers.zeros_init()
|
| 502 |
+
|
| 503 |
+
def get_encoder_attention_types(self) -> List[str]:
|
| 504 |
+
"""Get expanded per-layer encoder attention types.
|
| 505 |
+
|
| 506 |
+
If encoder_attention_types is None, returns [attention_mechanism] * encoder_nlayer.
|
| 507 |
+
If specified, expands short patterns by repetition (pattern length must divide layer count).
|
| 508 |
+
"""
|
| 509 |
+
return self._expand_attention_types(
|
| 510 |
+
self.encoder_attention_types,
|
| 511 |
+
self.encoder_nlayer,
|
| 512 |
+
"encoder_attention_types"
|
| 513 |
+
)
|
| 514 |
+
|
| 515 |
+
def get_decoder_attention_types(self) -> List[str]:
|
| 516 |
+
"""Get expanded per-layer decoder self-attention types."""
|
| 517 |
+
return self._expand_attention_types(
|
| 518 |
+
self.decoder_attention_types,
|
| 519 |
+
self.decoder_nlayer,
|
| 520 |
+
"decoder_attention_types"
|
| 521 |
+
)
|
| 522 |
+
|
| 523 |
+
def get_decoder_cross_attention_types(self) -> List[str]:
|
| 524 |
+
"""Get expanded per-layer decoder cross-attention types."""
|
| 525 |
+
return self._expand_attention_types(
|
| 526 |
+
self.decoder_cross_attention_types,
|
| 527 |
+
self.decoder_nlayer,
|
| 528 |
+
"decoder_cross_attention_types"
|
| 529 |
+
)
|
| 530 |
+
|
| 531 |
+
def _expand_attention_types(
|
| 532 |
+
self,
|
| 533 |
+
attention_types: Optional[List[str]],
|
| 534 |
+
num_layers: int,
|
| 535 |
+
field_name: str
|
| 536 |
+
) -> List[str]:
|
| 537 |
+
"""Expand short attention type patterns to full layer count.
|
| 538 |
+
|
| 539 |
+
Args:
|
| 540 |
+
attention_types: List of attention types or None
|
| 541 |
+
num_layers: Number of layers to expand to
|
| 542 |
+
field_name: Name of field for error messages
|
| 543 |
+
|
| 544 |
+
Returns:
|
| 545 |
+
List of attention types with length == num_layers
|
| 546 |
+
"""
|
| 547 |
+
if attention_types is None:
|
| 548 |
+
if self.attention_mechanism is None:
|
| 549 |
+
raise ValueError(
|
| 550 |
+
f"{field_name} is None and no global attention_mechanism fallback. "
|
| 551 |
+
f"This should have been caught by _validate_attention_config()."
|
| 552 |
+
)
|
| 553 |
+
return [self.attention_mechanism] * num_layers
|
| 554 |
+
|
| 555 |
+
if len(attention_types) == num_layers:
|
| 556 |
+
return attention_types
|
| 557 |
+
|
| 558 |
+
if num_layers % len(attention_types) != 0:
|
| 559 |
+
raise ValueError(
|
| 560 |
+
f"{field_name} length ({len(attention_types)}) must divide "
|
| 561 |
+
f"layer count ({num_layers}) evenly for pattern repetition"
|
| 562 |
+
)
|
| 563 |
+
|
| 564 |
+
repeats = num_layers // len(attention_types)
|
| 565 |
+
return attention_types * repeats
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:bbdf4f887074486d60cce4296cc42e4c00bbf3f3873ac562bf7c87e88269f9c9
|
| 3 |
+
size 755547592
|
modeling_han2han.py
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
spiece.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ca2b6bd503cd712bb538a5619a688dfee434b6a32a165635de4940e0a9fe22d4
|
| 3 |
+
size 912333
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"backend": "tokenizers",
|
| 3 |
+
"bos_token": "<s>",
|
| 4 |
+
"clean_up_tokenization_spaces": false,
|
| 5 |
+
"eos_token": "</s>",
|
| 6 |
+
"mask_token": "<mask>",
|
| 7 |
+
"model_input_names": [
|
| 8 |
+
"input_ids",
|
| 9 |
+
"attention_mask"
|
| 10 |
+
],
|
| 11 |
+
"model_max_length": 2048,
|
| 12 |
+
"pad_token": "<pad>",
|
| 13 |
+
"tokenizer_class": "TokenizersBackend",
|
| 14 |
+
"unk_token": "<unk>"
|
| 15 |
+
}
|