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
modeling_han2han: the subword tables follow the model's device on load; continuous batching withheld (serve falls back to sequential generation)
Browse files- modeling_han2han.py +24 -11
modeling_han2han.py
CHANGED
|
@@ -1066,18 +1066,16 @@ class Han2HanPreTrainedModel(PreTrainedModel):
|
|
| 1066 |
)
|
| 1067 |
|
| 1068 |
if safetensors_path is not None and os.path.exists(safetensors_path):
|
|
|
|
|
|
|
| 1069 |
with safe_open(safetensors_path, framework="pt") as f:
|
| 1070 |
-
|
| 1071 |
-
|
| 1072 |
-
|
| 1073 |
-
|
| 1074 |
-
|
| 1075 |
-
|
| 1076 |
-
|
| 1077 |
-
if "decoder.jbu" in f.keys() and not hasattr(base_model.decoder, 'jbu'):
|
| 1078 |
-
base_model.decoder.register_buffer('jbu', f.get_tensor("decoder.jbu"))
|
| 1079 |
-
if "decoder.cbu" in f.keys() and not hasattr(base_model.decoder, 'cbu'):
|
| 1080 |
-
base_model.decoder.register_buffer('cbu', f.get_tensor("decoder.cbu"))
|
| 1081 |
|
| 1082 |
if loading_info is not None:
|
| 1083 |
return model, loading_info
|
|
@@ -1582,6 +1580,21 @@ class Han2Han(Han2HanPreTrainedModel, GenerationMixin):
|
|
| 1582 |
self.encoder.jbu = self.decoder.jbu = torch.tensor(jamo_buckets.copy())
|
| 1583 |
self.register_buffer('jbu', self.encoder.jbu)
|
| 1584 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1585 |
def tie_weights(self, missing_keys=None, *args, **kwargs):
|
| 1586 |
"""Override HF's generic tie_weights.
|
| 1587 |
|
|
|
|
| 1066 |
)
|
| 1067 |
|
| 1068 |
if safetensors_path is not None and os.path.exists(safetensors_path):
|
| 1069 |
+
# the base class has already placed the weights (device_map, .to()), so each
|
| 1070 |
+
# table goes to the device of the embedding it is indexed beside
|
| 1071 |
with safe_open(safetensors_path, framework="pt") as f:
|
| 1072 |
+
for module, module_name in ((base_model.encoder, "encoder"), (base_model.decoder, "decoder")):
|
| 1073 |
+
for table_name in ("jbu", "cbu"):
|
| 1074 |
+
key = f"{module_name}.{table_name}"
|
| 1075 |
+
if key in f.keys() and not hasattr(module, table_name):
|
| 1076 |
+
module.register_buffer(
|
| 1077 |
+
table_name, f.get_tensor(key).to(module.wte.weight.device)
|
| 1078 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1079 |
|
| 1080 |
if loading_info is not None:
|
| 1081 |
return model, loading_info
|
|
|
|
| 1580 |
self.encoder.jbu = self.decoder.jbu = torch.tensor(jamo_buckets.copy())
|
| 1581 |
self.register_buffer('jbu', self.encoder.jbu)
|
| 1582 |
|
| 1583 |
+
@property
|
| 1584 |
+
def init_continuous_batching(self):
|
| 1585 |
+
"""Continuous batching is not available for this model.
|
| 1586 |
+
|
| 1587 |
+
`GenerationMixin` offers it to every model, but it drives a decoder-only forward
|
| 1588 |
+
pass over one packed sequence with a paged KV cache. Han2Han is an encoder-decoder
|
| 1589 |
+
with its own attention, so the attribute is withheld: `transformers serve
|
| 1590 |
+
--continuous-batching` then falls back to sequential generation with a warning
|
| 1591 |
+
instead of failing inside `forward()`.
|
| 1592 |
+
"""
|
| 1593 |
+
raise AttributeError(
|
| 1594 |
+
f"{type(self).__name__} does not support continuous batching (encoder-decoder model); "
|
| 1595 |
+
"use generate()."
|
| 1596 |
+
)
|
| 1597 |
+
|
| 1598 |
def tie_weights(self, missing_keys=None, *args, **kwargs):
|
| 1599 |
"""Override HF's generic tie_weights.
|
| 1600 |
|