|
13 | 13 | from typing import TYPE_CHECKING, Any |
14 | 14 |
|
15 | 15 | from transformers import ( |
| 16 | + AutoTokenizer, |
16 | 17 | CONFIG_MAPPING, |
17 | 18 | MODEL_MAPPING, |
18 | 19 | TOKENIZER_MAPPING, |
@@ -252,6 +253,54 @@ def from_backbone_class(self, BackboneClass: type[PreTrainedModel]) -> type[Ligh |
252 | 253 | class LightningIRTokenizerClassFactory(LightningIRClassFactory): |
253 | 254 | """Class factory for creating derived LightningIRTokenizer classes from HuggingFace tokenizer classes.""" |
254 | 255 |
|
| 256 | + @classmethod |
| 257 | + def _get_backbone_tokenizer_class( |
| 258 | + cls, |
| 259 | + backbone_config: PretrainedConfig, |
| 260 | + model_name_or_path: str | Path | None = None, |
| 261 | + use_fast: bool = True, |
| 262 | + ) -> type[PreTrainedTokenizerBase]: |
| 263 | + """Resolve a concrete tokenizer class for a backbone config. |
| 264 | +
|
| 265 | + Transformers v5 can contain configs that are not directly present in |
| 266 | + ``TOKENIZER_MAPPING`` even though tokenizer names are available via |
| 267 | + ``TOKENIZER_MAPPING_NAMES``. |
| 268 | + """ |
| 269 | + try: |
| 270 | + tokenizer_class_or_tuple = TOKENIZER_MAPPING[type(backbone_config)] |
| 271 | + return cls._resolve_tokenizer_class(tokenizer_class_or_tuple, use_fast=use_fast) |
| 272 | + except KeyError: |
| 273 | + model_type = backbone_config.model_type |
| 274 | + tokenizer_names = TOKENIZER_MAPPING_NAMES.get(model_type) |
| 275 | + if tokenizer_names is None: |
| 276 | + # Some newer model types are not wired into TOKENIZER_MAPPING_NAMES. |
| 277 | + # Fall back to AutoTokenizer for the concrete class. |
| 278 | + source = str(model_name_or_path or backbone_config.name_or_path) |
| 279 | + return AutoTokenizer.from_pretrained(source, use_fast=use_fast).__class__ |
| 280 | + |
| 281 | + module_name = model_type_to_module_name(model_type) |
| 282 | + module = importlib.import_module(f".{module_name}", "transformers.models") |
| 283 | + |
| 284 | + names: list[str] = [] |
| 285 | + if isinstance(tokenizer_names, tuple): |
| 286 | + slow_name, fast_name = tokenizer_names |
| 287 | + if use_fast and fast_name is not None: |
| 288 | + names.append(fast_name) |
| 289 | + if slow_name is not None: |
| 290 | + names.append(slow_name) |
| 291 | + if not use_fast and fast_name is not None: |
| 292 | + names.append(fast_name) |
| 293 | + else: |
| 294 | + names.append(tokenizer_names) |
| 295 | + |
| 296 | + for name in names: |
| 297 | + if hasattr(module, name): |
| 298 | + tokenizer_class = getattr(module, name) |
| 299 | + if isinstance(tokenizer_class, type): |
| 300 | + return tokenizer_class |
| 301 | + |
| 302 | + raise ValueError(f"Could not resolve tokenizer class for model_type '{model_type}'.") |
| 303 | + |
255 | 304 | @staticmethod |
256 | 305 | def _resolve_tokenizer_class( |
257 | 306 | tokenizer_class_or_tuple: type[PreTrainedTokenizerBase] | tuple[type[PreTrainedTokenizerBase] | None, ...], |
@@ -330,7 +379,9 @@ def from_pretrained( |
330 | 379 | type[LightningIRTokenizer]: Derived LightningIRTokenizer. |
331 | 380 | """ |
332 | 381 | backbone_config = self.get_backbone_config(model_name_or_path) |
333 | | - BackboneTokenizer = self._resolve_tokenizer_class(TOKENIZER_MAPPING[type(backbone_config)], use_fast=use_fast) |
| 382 | + BackboneTokenizer = self._get_backbone_tokenizer_class( |
| 383 | + backbone_config, model_name_or_path=model_name_or_path, use_fast=use_fast |
| 384 | + ) |
334 | 385 | DerivedLightningIRTokenizer = self.from_backbone_class(BackboneTokenizer) |
335 | 386 | return DerivedLightningIRTokenizer |
336 | 387 |
|
|
0 commit comments