| from pathlib import Path |
| import sys |
|
|
| _PROJECT_FILE = Path(__file__).resolve() |
| for _PROJECT_ROOT in _PROJECT_FILE.parents: |
| if (_PROJECT_ROOT / "model").is_dir(): |
| if str(_PROJECT_ROOT) not in sys.path: |
| sys.path.insert(0, str(_PROJECT_ROOT)) |
| break |
| |
|
|
| |
| |
|
|
| import torch |
| from fairscale.nn.data_parallel import FullyShardedDataParallel as FSDP |
| from fairscale.nn.wrap import enable_wrap, wrap |
|
|
| import model.esm as esm |
|
|
| |
| url = "tcp://localhost:23456" |
| torch.distributed.init_process_group(backend="nccl", init_method=url, world_size=1, rank=0) |
|
|
| |
| model_name = "esm2_t48_15B_UR50D" |
| model_data, regression_data = esm.pretrained._download_model_and_regression_data(model_name) |
|
|
| |
| fsdp_params = dict( |
| mixed_precision=True, |
| flatten_parameters=True, |
| state_dict_device=torch.device("cpu"), |
| cpu_offload=True, |
| ) |
| with enable_wrap(wrapper_cls=FSDP, **fsdp_params): |
| model, vocab = esm.pretrained.load_model_and_alphabet_core( |
| model_name, model_data, regression_data |
| ) |
| batch_converter = vocab.get_batch_converter() |
| model.eval() |
|
|
| |
| for name, child in model.named_children(): |
| if name == "layers": |
| for layer_name, layer in child.named_children(): |
| wrapped_layer = wrap(layer) |
| setattr(child, layer_name, wrapped_layer) |
| model = wrap(model) |
|
|
| data = [ |
| ("protein1", "MKTVRQERLKSIVRILERSKEPVSGAQLAEELSVSRQVIVQDIAYLRSLGYNIVATPRGYVLAGG"), |
| ("protein2", "KALTARQQEVFDLIRDHISQTGMPPTRAEIAQRLGFRSPNAAEEHLKALARKGVIEIVSGASRGIRLLQEE"), |
| ( |
| "protein2 with mask", |
| "KALTARQQEVFDLIRD<mask>ISQTGMPPTRAEIAQRLGFRSPNAAEEHLKALARKGVIEIVSGASRGIRLLQEE", |
| ), |
| ("protein3", "K A <mask> I S Q"), |
| ] |
|
|
| batch_labels, batch_strs, batch_tokens = batch_converter(data) |
| batch_tokens = batch_tokens.cuda() |
| with torch.no_grad(): |
| results = model(batch_tokens, repr_layers=[48], return_contacts=True) |
| print(results) |
|
|