Model pretrained on protein and antibody sequences using a masked language modeling (MLM) objective. It was introduced in the paper
Large scale paired antibody language models.
The model is finetuned from IgT5-unpaired using paired antibody sequences from the
Observed Antibody Space.
1from transformers import T5EncoderModel, T5Tokenizer
2
3tokeniser = T5Tokenizer.from_pretrained("Exscientia/IgT5", do_lower_case=False)
4model = T5EncoderModel.from_pretrained("Exscientia/IgT5")
1# heavy chain sequences
2sequences_heavy = [
3 "VQLAQSGSELRKPGASVKVSCDTSGHSFTSNAIHWVRQAPGQGLEWMGWINTDTGTPTYAQGFTGRFVFSLDTSARTAYLQISSLKADDTAVFYCARERDYSDYFFDYWGQGTLVTVSS",
4 "QVQLVESGGGVVQPGRSLRLSCAASGFTFSNYAMYWVRQAPGKGLEWVAVISYDGSNKYYADSVKGRFTISRDNSKNTLYLQMNSLRTEDTAVYYCASGSDYGDYLLVYWGQGTLVTVSS"
5]
6
7# light chain sequences
8sequences_light = [
9 "EVVMTQSPASLSVSPGERATLSCRARASLGISTDLAWYQQRPGQAPRLLIYGASTRATGIPARFSGSGSGTEFTLTISSLQSEDSAVYYCQQYSNWPLTFGGGTKVEIK",
10 "ALTQPASVSGSPGQSITISCTGTSSDVGGYNYVSWYQQHPGKAPKLMIYDVSKRPSGVSNRFSGSKSGNTASLTISGLQSEDEADYYCNSLTSISTWVFGGGTKLTVL"
11]
12
13# The tokeniser expects input of the form ["V Q ... S S </s> E V ... I K", ...]
14paired_sequences = []
15for sequence_heavy, sequence_light in zip(sequences_heavy, sequences_light):
16 paired_sequences.append(' '.join(sequence_heavy)+' </s> '+' '.join(sequence_light))
17
18tokens = tokeniser.batch_encode_plus(
19 paired_sequences,
20 add_special_tokens=True,
21 pad_to_max_length=True,
22 return_tensors="pt",
23 return_special_tokens_mask=True
24)
1output = model(
2 input_ids=tokens['input_ids'],
3 attention_mask=tokens['attention_mask']
4)
5
6residue_embeddings = output.last_hidden_state
To obtain a sequence representation, the residue tokens can be averaged over like so
1import torch
2
3# mask special tokens before summing over embeddings
4residue_embeddings[tokens["special_tokens_mask"] == 1] = 0
5sequence_embeddings_sum = residue_embeddings.sum(1)
6
7# average embedding by dividing sum by sequence lengths
8sequence_lengths = torch.sum(tokens["special_tokens_mask"] == 0, dim=1)
9sequence_embeddings = sequence_embeddings_sum / sequence_lengths.unsqueeze(1)