Views
No views yet
1import esm
2import torch
3import reverse_distillation
4
5# Load ESM-2 model and the reverse distillation version
6esm2_model, alphabet = esm.pretrained.esm2_t33_650M_UR50D()
7rd_model, alphabet = reverse_distillation.pretrained.esm2_rd_650M()
8
9batch_converter = alphabet.get_batch_converter()
10esm2_model.eval() # disables dropout for deterministic results
11rd_model.eval() # disables dropout for deterministic results
12
13# Prepare data
14data = [
15 ("protein1", "MKTVRQERLKSIVRILERSKEPVSGAQLAEELSVSRQVIVQDIAYLRSLGYNIVATPRGYVLAGG"),
16 ("protein2", "KALTARQQEVFDLIRDHISQTGMPPTRAEIAQRLGFRSPNAAEEHLKALARKGVIEIVSGASRGIRLLQEE"),
17]
18batch_labels, batch_strs, batch_tokens = batch_converter(data)
19batch_lens = (batch_tokens != alphabet.padding_idx).sum(1)
20
21# Extract per-residue representations
22with torch.no_grad():
23 results_esm = esm2_model(batch_tokens, repr_layers=[33], return_contacts=True)
24 results_rd = rd_model(batch_tokens)
25
26esm_token_representations = results_esm["representations"][33]
27rd_token_representations = results_rd["representations"]["650M"]
28
29# Generate per-sequence representations via averaging
30for i, tokens_len in enumerate(batch_lens):
31 print(f"esm representation size: {esm_token_representations[i, 1 : tokens_len - 1].size()}")
32 print(f"rd representation size: {rd_token_representations[i, 1 : tokens_len - 1].size()}")1@inproceedings{catrina2026reverse,
2 title = {Reverse Distillation: Consistently Scaling Protein Language Model Representations},
3 author = {Catrina, Darius and Bepler, Christian and Sledzieski, Samuel and Singh, Rohit},
4 booktitle = {International Conference on Learning Representations},
5 year = {2026}
6}