Views
No views yet
1from transformers import T5ForConditionalGeneration, T5Tokenizer, T5EncoderModel
2import torch
3
4# Random sequence from uniprot, most likely Ankh3 saw it during pre-training.
5sequence = "MDTAYPREDTRAPTPSKAGAHTALTLGAPHPPPRDHLIWSVFSTLYLNLCCLGFLALAYSIKARDQKVVGDLEAARRFGSKAKCYNILAAMWTLVPPLLLLGLVVTGALHLARLAKDSAAFFSTKFDDADYD"
6
7ckpt = "ElnaggarLab/ankh3-xl"
8
9# Make sure that you must use `T5Tokenizer` not `AutoTokenizer`.
10tokenizer = T5Tokenizer.from_pretrained(ckpt)
11
12# To use the encoder representation using the NLU prefix:
13encoder_model = T5EncoderModel.from_pretrained(ckpt).eval()
14
15
16# For extracting embeddings, consider trying the '[S2S]' prefix.
17# Since this prefix was specifically used to denote sequence completion
18# during the model's pre-training, its use can sometimes
19# lead to improved embedding quality.
20
21nlu_sequence = "[NLU]" + sequence
22encoded_nlu_sequence = tokenizer(nlu_sequence, add_special_tokens=True, return_tensors="pt", is_split_into_words=False)
23
24with torch.no_grad():
25 embedding = encoder_model(**encoded_nlu_sequence)1from transformers import T5ForConditionalGeneration, T5Tokenizer
2from transformers.generation import GenerationConfig
3import torch
4
5sequence = "MDTAYPREDTRAPTPSKAGAHTALTLGAPHPPPRDHLIWSVFSTLYLNLCCLGFLALAYSIKARDQKVVGDLEAARRFGSKAKCYNILAAMWTLVPPLLLLGLVVTGALHLARLAKDSAAFFSTKFDDADYD"
6
7ckpt = "ElnaggarLab/ankh3-xl"
8tokenizer = T5Tokenizer.from_pretrained(ckpt)
9# To use the sequence to sequence task using the S2S prefix:
10model = T5ForConditionalGeneration.from_pretrained(ckpt).eval()
11
12
13half_length = int(len(sequence) * 0.5)
14s2s_sequence = "[S2S]" + sequence[:half_length]
15encoded_s2s_sequence = tokenizer(s2s_sequence, add_special_tokens=True, return_tensors="pt", is_split_into_words=False)
16# + 1 to account for the start of sequence token.
17gen_config = GenerationConfig(min_length=half_length + 1, max_length=half_length + 1, do_sample=False, num_beams=1)
18generated_sequence = model.generate(encoded_s2s_sequence["input_ids"], gen_config, )
19predicted_sequence = sequence[:half_length] + tokenizer.batch_decode(generated_sequence)[0]