(name,
encoder_only=False,
decoder_only=False,
return_tokenizer=False,
tokenizer_kwargs={},
dtype=torch.float32,
device='cpu',
**kwargs)
| 413 | |
| 414 | |
| 415 | def _t5(name, |
| 416 | encoder_only=False, |
| 417 | decoder_only=False, |
| 418 | return_tokenizer=False, |
| 419 | tokenizer_kwargs={}, |
| 420 | dtype=torch.float32, |
| 421 | device='cpu', |
| 422 | **kwargs): |
| 423 | # sanity check |
| 424 | assert not (encoder_only and decoder_only) |
| 425 | |
| 426 | # params |
| 427 | if encoder_only: |
| 428 | model_cls = T5Encoder |
| 429 | kwargs['vocab'] = kwargs.pop('vocab_size') |
| 430 | kwargs['num_layers'] = kwargs.pop('encoder_layers') |
| 431 | _ = kwargs.pop('decoder_layers') |
| 432 | elif decoder_only: |
| 433 | model_cls = T5Decoder |
| 434 | kwargs['vocab'] = kwargs.pop('vocab_size') |
| 435 | kwargs['num_layers'] = kwargs.pop('decoder_layers') |
| 436 | _ = kwargs.pop('encoder_layers') |
| 437 | else: |
| 438 | model_cls = T5Model |
| 439 | |
| 440 | # init model |
| 441 | with torch.device(device): |
| 442 | model = model_cls(**kwargs) |
| 443 | |
| 444 | # set device |
| 445 | model = model.to(dtype=dtype, device=device) |
| 446 | |
| 447 | # init tokenizer |
| 448 | if return_tokenizer: |
| 449 | from .tokenizers import HuggingfaceTokenizer |
| 450 | tokenizer = HuggingfaceTokenizer(f'google/{name}', **tokenizer_kwargs) |
| 451 | return model, tokenizer |
| 452 | else: |
| 453 | return model |
| 454 | |
| 455 | |
| 456 | def umt5_xxl(**kwargs): |
no test coverage detected