cadazar commited on
Commit
bfd235a
·
verified ·
1 Parent(s): e77af7c

Han2Han pre-trained checkpoint (uniform average of the last five checkpoints, steps 6347 to 6725)

Browse files
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
+ }