MULTI-evolve / scripts /p2_propose.py
anzhi2710gmailcom's picture
Upload folder using huggingface_hub
6f1e670 verified
Raw
History Blame Contribute Delete
8.02 kB
#!/public/home/scnb9biwet/.conda/envs/model_bio/bin/python
"""
Script to propose mutations using trained multievolve models.
Modified to load local models instead of using wandb.
"""
import os
os.environ["WANDB_MODE"] = "disabled"
os.environ["WANDB_SILENT"] = "true"
import wandb
import argparse
import pandas as pd
import numpy as np
from Bio import SeqIO
import matplotlib
matplotlib.use('Agg')
import glob
import re
import torch
from model.splitters import *
from model.featurizers import *
from model.predictors import *
from model.proposers import *
def parse_args():
"""Parse command line arguments."""
parser = argparse.ArgumentParser(description='Propose mutations using trained models')
parser.add_argument(
'--experiment-name',
required=True,
help='Name of experiment'
)
parser.add_argument(
'--protein-name',
required=True,
help='Name of protein'
)
parser.add_argument(
'--wt-files',
required=True,
help='Comma separated list of paths to the wildtype FASTA files'
)
parser.add_argument(
'--training-dataset',
required=True,
help='Path to training dataset CSV'
)
parser.add_argument(
'--mutation-pool',
required=True,
help='Path to mutation pool CSV'
)
parser.add_argument(
'--top-muts-per-load',
type=int,
default=3,
help='Number of top mutations to select per load (default: 3)'
)
parser.add_argument(
'--export-name',
required=True,
help='Name for export files'
)
args = parser.parse_args()
args.wt_files = [f.strip() for f in args.wt_files.split(',')]
return args
def main():
"""Main function."""
# Parse command line arguments
args = parse_args()
# Define variables from args
experiment_name = args.experiment_name
protein_name = args.protein_name
wt_files = args.wt_files
training_dataset_fname = args.training_dataset
mutation_pool_fname = args.mutation_pool
top_muts_per_load = args.top_muts_per_load
export_name = args.export_name
# Processed variables
mutation_pool = pd.read_csv(mutation_pool_fname, header=None).values.flatten().tolist()
wt_seq = "".join([str(SeqIO.read(wt_file, "fasta").seq.upper()) for wt_file in wt_files])
# 手动指定最佳超参数(来自 fcn_test_sweep.yaml 和训练设置)
bs = 32
lr = 0.0001
hidden = 100
layers = 1
print(bs, lr, hidden, layers)
# 配置模型
config = {
'layer_size': hidden,
'num_layers' : layers,
'learning_rate': lr,
'batch_size': bs,
'optimizer': 'adam',
'epochs': 300
}
# 初始化 splits(与训练时一致,5 折)
split = KFoldProteinSplitter(protein_name, training_dataset_fname, wt_files,
csv_has_header=True, use_cache=True, y_scaling=True, val_split=0.15)
splits = split.generate_splits(n_splits=5)
# 初始化 feature
feature = OneHotFeaturizer(protein=protein_name, use_cache=True)
# 加载已有模型
# 从分裂对象中获取 dataset_dir 和 dataset_name,动态构造模型目录
dataset_dir = splits[0].file_attrs['dataset_dir']
dataset_name = splits[0].file_attrs['dataset_name']
model_dir = os.path.join(dataset_dir, 'model_cache', dataset_name, 'objects')
#model_dir = os.path.join(splits[0].file_attrs['model_dir'], 'objects')
model_files = glob.glob(os.path.join(model_dir, 'split_by_kfold-*.pth'))
# 按 fold 编号排序
model_files.sort(key=lambda x: int(re.search(r'split_by_kfold-(\d+)_', os.path.basename(x)).group(1)))
print(f"Found {len(model_files)} model files in {model_dir}")
models = []
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
for i, split in enumerate(splits):
model = Fcn(split, feature, config=config, use_cache=True)
model.load_state_dict(torch.load(model_files[i], map_location=device, weights_only=True))
model.to(device)
model.eval()
models.append(model)
print(f"Loaded model from {model_files[i]}")
print("Proposing mutations...")
# 初始化 proposer 并评估提案
proposer = CombinatorialProposer(
start_seq=wt_seq,
models=models,
trust_radius=11,
num_seeds=-1, # evaluate all seeds
mutation_pool=mutation_pool)
proposer.propose(output_df=False)
proposer.evaluate_proposals()
proposer.save_proposals(f'{experiment_name}_proposals_all')
# 获取每个突变负荷的前 N 个变体
df = proposer.proposals
df_ls = []
for num_mut in range(3, 11, 1):
subset = df[df['num_muts'] == num_mut].copy()
subset.sort_values(by='average', ascending=False, inplace=True)
top_subset = subset.head(top_muts_per_load).copy()
df_ls.append(top_subset)
top_df = pd.concat(df_ls, ignore_index=True)
# 导出结果
print('Saving all proposals...')
top_df.to_csv(os.path.join(splits[0].file_attrs['dataset_dir'], 'proposers/results',
f'{experiment_name}_proposals_top_{top_muts_per_load}.csv'), index=False)
top_df[['Mut_string']].to_csv(os.path.join(splits[0].file_attrs['dataset_dir'], f'{export_name}.csv'),
index=False, header=None)
# 多链蛋白处理函数(原代码保留,未修改)
def reverse_multichain_mutations(mut_strings, chain_lengths):
cumulative_lengths = [sum(chain_lengths[:i]) for i in range(len(chain_lengths))]
mutation_map = {}
for mut_string in mut_strings:
mutations = mut_string.split('/')
chain_mutations = {i: [] for i in range(len(chain_lengths))}
for mut in mutations:
position = int(mut[1:-1])
wt_aa = mut[0]
mut_aa = mut[-1]
for chain_idx, start_pos in enumerate(cumulative_lengths):
if position <= cumulative_lengths[chain_idx + 1] if chain_idx + 1 < len(cumulative_lengths) else float('inf'):
chain_pos = position - start_pos
chain_mutations[chain_idx].append(f"{wt_aa}{chain_pos}{mut_aa}")
break
mutation_map[mut_string] = chain_mutations
return mutation_map
def mutation_map_to_df(mutation_map):
rows = []
for mut_string, chain_muts in mutation_map.items():
row = {'Mut_string': mut_string}
for chain_idx, mutations in chain_muts.items():
row[f'chain_{chain_idx + 1}'] = '/'.join(mutations) if mutations else ''
rows.append(row)
df = pd.DataFrame(rows)
chain_cols = [col for col in df.columns if col.startswith('chain_')]
df = df[['Mut_string'] + sorted(chain_cols)]
return df
if len(wt_files) > 1:
mutations = top_df['Mut_string'].values.tolist()
chain_lens = splits[0].wt_seq_lens
dict_mutations = reverse_multichain_mutations(mutations, chain_lens)
df_mutations = mutation_map_to_df(dict_mutations)
top_df = pd.merge(top_df, df_mutations, on='Mut_string', how='left')
top_df.to_csv(os.path.join(splits[0].file_attrs['dataset_dir'], 'proposers/results',
f'{experiment_name}_proposals_top_{top_muts_per_load}.csv'), index=False)
for col in df_mutations.columns[1:]:
mutations = set(df_mutations[col].tolist())
if '' in mutations:
mutations.remove('')
df_mutations_col = pd.DataFrame(mutations, columns=[col])
df_mutations_col.to_csv(os.path.join(splits[0].file_attrs['dataset_dir'],
f'{export_name}_{col}_mutants.csv'), index=False, header=None)
if __name__ == "__main__":
main()