MCPcopy Create free account
hub / github.com/MotrixLab/ViMoGen / _t5

Function _t5

models/transformer/wan/modules/t5.py:415–453  ·  view source on GitHub ↗
(name,
        encoder_only=False,
        decoder_only=False,
        return_tokenizer=False,
        tokenizer_kwargs={},
        dtype=torch.float32,
        device='cpu',
        **kwargs)

Source from the content-addressed store, hash-verified

413
414
415def _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
456def umt5_xxl(**kwargs):

Callers 1

umt5_xxlFunction · 0.85

Calls 1

Tested by

no test coverage detected