OneScience-Group/SurfDock
025
1"""2caoduanhua : we should to implemented a parapllel version of evaluate.py for a large dataset3"""4 5import copy6import os7import sys8 9SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))10PROJECT_DIR = os.path.dirname(SCRIPT_DIR)11MODEL_DIR = os.path.join(PROJECT_DIR, "model")12if MODEL_DIR not in sys.path:13 sys.path.insert(0, MODEL_DIR)14 15import torch16import time17from argparse import ArgumentParser, Namespace, FileType18from datetime import datetime19from functools import partial20import numpy as np21import gc22import pandas as pd23import wandb24import glob25from rdkit import RDLogger26from rdkit.Chem import RemoveHs27from datasets.process_mols import write_mol_with_coords28from torch_geometric.loader import DataLoader29from datasets.pdbbind import PDBBind, read_mol,read_abs_file_mol30from utils.diffusion_utils import t_to_sigma as t_to_sigma_compl, get_t_schedule31from utils.sampling import randomize_position, sampling,inferenceFFOptimize32from utils.utils import get_symmetry_rmsd, remove_all_hs33from score_in_place_dataset.score_dataset import ScreenDataset34from utils.utils import get_model, ExponentialMovingAverage35from utils.visualise import PDBFile36from tqdm import tqdm37from collections import defaultdict38from packaging import version39import warnings40warnings.filterwarnings("ignore", category=UserWarning, module="torch.jit")41RDLogger.DisableLog('rdApp.*')42import yaml43from loguru import logger44 45cache_name = datetime.now().strftime('date%d-%m_time%H-%M-%S.%f')46parser = ArgumentParser()47 48parser.add_argument('--config', type=FileType(mode='r'), default=None)49parser.add_argument('--data_csv', type=str, default='~/Screen_dataset/dataset/DEKOIS2.csv', help='Path to folder with dataset for score in place')50parser.add_argument('--model_dir', type=str, default=None, help='Path to folder with trained score model and hyperparameters')51parser.add_argument('--ckpt', type=str, default=None, help='Checkpoint to use inside the folder')52parser.add_argument('--confidence_model_dir', type=str, default=None, help='Path to folder with trained confidence model and hyperparameters')53parser.add_argument('--confidence_ckpt', type=str, default=None, help='Checkpoint to use inside the folder')54# save docking result or not55parser.add_argument('--save_docking_result', action='store_true', default=False, help='Whether to save docking result')56# put ligand to pocket center57parser.add_argument('--ligand_to_pocket_center', action='store_true', default=False, help='Whether to put ligand on pocket center')58parser.add_argument('--keep_input_pose', action='store_false', default=False, help='Whether keep original input pose')59parser.add_argument('--use_noise_to_rank', action='store_true', default=False, help='Whether to run the probability flow ODE')60parser.add_argument('--num_cpu', type=int, default=None, help='if this is a number instead of none, the max number of cpus used by torch will be set to this.')61parser.add_argument('--run_name', type=str, default='test_ns_48_nv_10_layer_62023-06-25_07-54-08_model', help='')62parser.add_argument('--project', type=str, default='ligbind_inf_test_mdn', help='')63parser.add_argument('--surface_path', type=str, default='~/PDBBind_processed_8A_surface/', help='test dataset surface path')64parser.add_argument('--esm_embeddings_path', type=str, default='~/PDBBIND/esm_embedding/esm_embedding_pocket_for_train/esm2_3billion_embeddings.pt', help='test dataset esmbedding path')65parser.add_argument('--out_dir', type=str, default='~/test_workdir/mdn_result_40', help='Where to save results to')66parser.add_argument('--batch_size', type=int, default=40, help='Number of poses to sample in parallel we recommand set number = batch_size_molecule*samples_per_complex')67parser.add_argument('--batch_size_molecule', type=int, default=1, help='Number of molecul to sample in parallel')68parser.add_argument('--cache_path', type=str, default='~/PDBBIND/cache_PDBBIND_pocket_8A', help='Folder from where to load/restore cached dataset')69parser.add_argument('--data_dir', type=str, default='~/PDBBIND/PDBBind_pocket_8A/', help='Folder containing original structures')70parser.add_argument('--split_path', type=str, default='~/data/splits/timesplit_test', help='Path of file defining the split')71parser.add_argument('--no_overlap_names_path', type=str, default='~/data/splits/timesplit_test_no_rec_overlap', help='Path text file with the folder names in the test set that have no receptor overlap with the train set')72parser.add_argument('--no_model', action='store_true', default=False, help='Whether to return seed conformer without running model')73parser.add_argument('--no_random', action='store_true', default=False, help='Whether to add randomness in diffusion steps')74parser.add_argument('--no_final_step_noise', action='store_true', default=False, help='Whether to add noise after the final step')75parser.add_argument('--ode', action='store_true', default=False, help='Whether to run the probability flow ODE')76parser.add_argument('--wandb', action='store_true', default=False, help='')77parser.add_argument('--wandb_dir', type=str, default='~/test_workdir', help='Folder in which to save wandb logs')78parser.add_argument('--inference_steps', type=int, default=20, help='Number of denoising steps')79parser.add_argument('--limit_complexes', type=int, default=0, help='Limit to the number of complexes')80parser.add_argument('--num_workers', type=int, default=1, help='Number of workers for dataset creation')81parser.add_argument('--num_process', type=int, default=20, help='Number of parallel workers for minimized.')82parser.add_argument('--tqdm', action='store_true', default=False, help='Whether to show progress bar')83parser.add_argument('--save_visualisation', action='store_true', default=False, help='Whether to save visualizations')84parser.add_argument('--samples_per_complex', type=int, default=40, help='Number of poses to sample for each complex')85parser.add_argument('--save_docking_result_number', type=int, default=1, help='Number of poses to save in disk for each complex')86parser.add_argument('--actual_steps', type=int, default=None, help='')87parser.add_argument('--inference_mode', default='Screen', help='inference mode',choices=['Screen','evaluate'])88parser.add_argument('--head_index', type=int, default=0, help='the head index to start inference,this optinal to inference use multi-GPU every GPU minimized a part of csv file ')89parser.add_argument('--tail_index', type=int, default=-1, help='the tail index to start inference,this optinal to inference use multi-GPU every GPU minimized a part of csv file')90parser.add_argument('--ligandsMaxAtoms', type=int, default=80, help='the max number of atoms in ligand')91parser.add_argument('--random_seed', type=int, default=42, help='random seed')92# force_minimized param93parser.add_argument('--force_optimize', action='store_true', default=False, help='')94parser.add_argument('--mdn_dist_threshold_test', type=float, default=3.0, help='mdn_dist_threshold_test')95args = parser.parse_args()96nowtime = datetime.now().strftime('%Y-%m-%d')97log_file_flag = '-'.join(args.project.split('/'))98logger.add(f'{os.path.dirname(args.out_dir)}/log-inference-{log_file_flag}-{nowtime}.log', rotation="500MB")99logger.info('Runing inference script in path: {}',os.getcwd())100logger.info('Runing inference with args: {}',args)101 102def main_function():103 if accelerator.is_local_main_process:104 if args.wandb:105 wandb.login(key = 'yourkey')106 run = wandb.init(107 entity='SurfDock',108 settings=wandb.Settings(start_method="fork"),109 project=args.project,110 name=args.run_name,111 dir = args.wandb_dir,112 config=args113 )114 if args.config:115 config_dict = yaml.load(args.config, Loader=yaml.FullLoader)116 arg_dict = args.__dict__117 for key, value in config_dict.items():118 if isinstance(value, list):119 for v in value:120 arg_dict[key].append(v)121 else:122 arg_dict[key] = value123 if args.out_dir is None: args.out_dir = f'inference_out_dir_not_specified/{args.run_name}'124 os.makedirs(args.out_dir, exist_ok=True)125 with open(f'{args.model_dir}/model_parameters.yml') as f:126 score_model_args = Namespace(**yaml.full_load(f))127 128 129 if args.confidence_model_dir is not None:130 with open(f'{args.confidence_model_dir}/model_parameters.yml') as f:131 confidence_args = Namespace(**yaml.full_load(f))132 # 133 confidence_args.transfer_weights = False134 confidence_args.use_original_model_cache = True135 confidence_args.original_model_dir = None136 confidence_args.mdn_dist_threshold_test = args.mdn_dist_threshold_test if args.mdn_dist_threshold_test is not None else 5.0137 if not hasattr(confidence_args,'mdn_dist_threshold_train'):138 confidence_args.mdn_dist_threshold_train =7.0139 140 if args.confidence_model_dir is not None:141 if not (confidence_args.use_original_model_cache or confidence_args.transfer_weights):142 # if the confidence model uses the same type of data as the original model then we do not need this dataset and can just use the complexes143 logger.info('HAPPENING | confidence model uses different type of graphs than the score model. Loading (or creating if not existing) the data for the confidence model now.')144 confidence_test_dataset = PDBBind(transform=None, root=args.data_dir, limit_complexes=args.limit_complexes,145 receptor_radius=confidence_args.receptor_radius,146 cache_path=args.cache_path, split_path=args.split_path,147 remove_hs=confidence_args.remove_hs, max_lig_size=None, c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors,148 matching=not confidence_args.no_torsion, keep_original=True,149 popsize=confidence_args.matching_popsize,150 maxiter=confidence_args.matching_maxiter,151 all_atoms=confidence_args.all_atoms,152 atom_radius=confidence_args.atom_radius,153 atom_max_neighbors=confidence_args.atom_max_neighbors,154 esm_embeddings_path= args.esm_embeddings_path, require_ligand=True,155 num_workers=args.num_workers,surface_path = args.surface_path)156 confidence_complex_dict = {d.name: d for d in confidence_test_dataset}157 158 t_to_sigma = partial(t_to_sigma_compl, args=score_model_args)159 160 if not args.no_model:161 model = get_model(score_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,model_type = score_model_args.model_type)162 state_dict = torch.load(f'{args.model_dir}/{args.ckpt}', map_location=torch.device('cpu'))163 if args.ckpt == 'last_model.pt':164 model_state_dict = state_dict['model']165 ema_weights_state = state_dict['ema_weights']166 model.load_state_dict(model_state_dict, strict=True)167 ema_weights = ExponentialMovingAverage(model.parameters(), decay=score_model_args.ema_rate)168 ema_weights.load_state_dict(ema_weights_state, device=device)169 ema_weights.copy_to(model.parameters())170 else:171 model.load_state_dict(state_dict, strict=False)172 model = model.to(device)173 model.eval()174 logger.info('loaded model weight for score model')175 if args.confidence_model_dir is not None:176 if confidence_args.transfer_weights:177 with open(f'{confidence_args.original_model_dir}/model_parameters.yml') as f:178 confidence_model_args = Namespace(**yaml.full_load(f))179 else:180 confidence_model_args = confidence_args181 182 confidence_model = get_model(confidence_model_args, device, t_to_sigma=t_to_sigma, no_parallel=True,183 model_type = confidence_model_args.model_type)184 state_dict = torch.load(f'{args.confidence_model_dir}/{args.confidence_ckpt}', map_location=torch.device('cpu'))185 confidence_model.load_state_dict(state_dict, strict=True)186 confidence_model = confidence_model.to(device)187 confidence_model.eval()188 else:189 confidence_model = None190 confidence_args = None191 confidence_model_args = None192 193 194 tr_schedule = get_t_schedule(inference_steps=args.inference_steps)195 rot_schedule = tr_schedule196 tor_schedule = tr_schedule197 logger.info('t schedule:{}',tr_schedule)198 logger.info('Loading data ...........')199 200 """201 Load data from csv file to get the path of pocket,ligand,ref_ligand,surface202 203 """204 df = pd.read_csv(args.data_csv)[args.head_index:args.tail_index]205 protein_paths = df['protein_path'].tolist()206 pocket_paths = df['pocket_path'].tolist()207 ligands_paths = df['ligand_path'].tolist()208 ref_ligands = df['ref_ligand'].tolist()209 surface_paths = df['protein_surface'].tolist()210 if 'pocket_center' in df.columns:211 pocket_centers = df['pocket_center'].tolist()212 new_pocket_centers = []213 for center in pocket_centers:214 x = center.split(',')[0]215 y = center.split(',')[1]216 z = center.split(',')[2]217 new_pocket_centers.append(np.array([(float(x),float(y),float(z))]))218 pocket_centers = new_pocket_centers219 else:220 pocket_centers = [None]*len(protein_paths)221 222 esm_embeddings_dict = torch.load(args.esm_embeddings_path)223 confidence_list = []224 confidence_names = []225 sdf_names = []226 pocket_path_list =[]227 228 failures = 0229 N = args.samples_per_complex230 all_molecules = 0231 pbar = tqdm(zip(pocket_paths,ligands_paths,ref_ligands,surface_paths,protein_paths,pocket_centers),total=len(pocket_paths))232 start_time = time.time()233 for pocket_path,ligands_path,ref_ligand,surface_path,protein_path,pocket_center in pbar:234 in_loop_start_time = time.time()235 236 try:237 238 dirname = os.path.splitext(pocket_path.split('/')[-1])[0] + '_'+ os.path.splitext(ligands_path.split('/')[-1])[0] 239 write_dir = os.path.join(args.out_dir,'SurfDock_docking_result',dirname)#f'{args.out_dir}/SurfDock_docking_result/{dirname}'240 os.makedirs(write_dir, exist_ok=True)241 242 esm_embeddings = copy.deepcopy(esm_embeddings_dict[os.path.splitext(os.path.basename(pocket_path))[0]])243 244 test_dataset = ScreenDataset(pocket_path,ligands_path,ref_ligand,surface_path,pocket_center,transform=None,245 receptor_radius=confidence_args.receptor_radius,246 cache_path=None, split_path=None,247 remove_hs=confidence_args.remove_hs, max_lig_size=None,248 c_alpha_max_neighbors=confidence_args.c_alpha_max_neighbors,249 matching= False, keep_original=True,250 popsize=confidence_args.matching_popsize,251 maxiter=confidence_args.matching_maxiter,252 all_atoms=confidence_args.all_atoms,253 atom_radius=confidence_args.atom_radius,254 atom_max_neighbors=confidence_args.atom_max_neighbors,255 esm_embeddings=esm_embeddings,256 require_ligand=False,257 num_workers=args.num_workers,258 keep_input_pose = args.keep_input_pose,259 save_dir = write_dir,260 inference_mode = args.inference_mode,261 ligandsMaxAtoms=args.ligandsMaxAtoms)262 test_sample_num = len(test_dataset)263 all_molecules += test_sample_num264 test_loader = DataLoader(dataset=test_dataset, batch_size=args.batch_size_molecule, shuffle=False)265 if test_sample_num == 0:266 logger.error('No complexes need to be docking (skip before done or some errors) in {}', pocket_path)267 continue268 # test_loader= accelerator.prepare(test_loader)269 logger.info('Protein {} Size of test dataset: {}',os.path.splitext(os.path.basename(pocket_path))[0], test_sample_num)270 ##### use torch.__version__orch complie to speed up the process ###271 # if version.parse(torch.__version__.split('+')[0])> version.parse("2.0"):272 # model = torch.compile(model)273 # confidence_model = torch.compile(confidence_model)274 # logger.info('Your are using torch version={} , so SurfDock will use torch.compile to complie model and confidence model',torch.__version__)275 #########################################################276 model = accelerator.prepare(model)277 test_loader= accelerator.prepare(test_loader)278 confidence_model = accelerator.prepare(confidence_model)279 """280 Start sampling conformers by SurfDock281 """282 for idx, orig_complex_graph in tqdm(enumerate(test_loader),total = len(test_loader),disable= not accelerator.is_local_main_process):283 284 try:285 if 'ligand' not in orig_complex_graph.node_types:286 logger.error('some error failed for conformer generate in rdkit: idx in batch graph: {}, ligand_path: {}',idx,ligands_path)287 continue288 orig_complex_graph_list = orig_complex_graph.to_data_list()289 # add protein pocket information for minimized stage 290 for temp_graph in orig_complex_graph_list:291 temp_graph['protein_path'] = protein_path292 temp_graph['pocket_path'] = pocket_path293 294 success = 0295 sample_count_failed = 0296 data_list = []297 # object 298 data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]299 while not success: # keep trying in case of failure (sometimes stochastic)300 301 try:302 303 # data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]304 success = 1305 randomize_position(data_list, score_model_args.no_torsion, args.no_random, score_model_args.tr_sigma_max,ligand_to_pocket_center = args.ligand_to_pocket_center)306 pdb = None307 if args.save_visualisation:308 visualization_list = []309 for idx, graph in enumerate(data_list):310 # raw pose311 lig = read_mol(args.data_dir, graph['name'][0], remove_hs=score_model_args.remove_hs)312 pdb = PDBFile(lig)313 pdb.add(lig, 0, 0)314 # pose rdkit matching315 orig_complex_count = idx//N316 317 pdb.add((orig_complex_graph_list[orig_complex_count]['ligand'].pos + orig_complex_graph_list[orig_complex_count].original_center).detach().cpu(), 1, 0)318 # random rdkit matching319 pdb.add((graph['ligand'].pos + (graph.original_center).detach().cpu()), part=1, order=1)320 visualization_list.append(pdb)321 else:322 visualization_list = None323 324 if not args.no_model:325 326 confidence_data_list = None327 328 data_list, confidence = sampling(input_data_list=data_list, model=model,329 inference_steps=args.actual_steps if args.actual_steps is not None else args.inference_steps,330 tr_schedule=tr_schedule, rot_schedule=rot_schedule,331 tor_schedule=tor_schedule,332 device=device, t_to_sigma=t_to_sigma, model_args=score_model_args,333 no_random=args.no_random,334 ode=args.ode, visualization_list=visualization_list,335 confidence_model=confidence_model,336 confidence_data_list=confidence_data_list,337 confidence_model_args=confidence_model_args,338 batch_size=args.batch_size,339 no_final_step_noise=args.no_final_step_noise,args = args)340 accelerator.wait_for_everyone()341 342 confidence = confidence.cpu().detach().numpy()343 344 # save confidence345 confidence_list += confidence.tolist()346 for _ in range(len(orig_complex_graph_list)):347 348 confidence_names.extend([orig_complex_graph_list[_]['name']]*N)349 pocket_path_list.extend([os.path.basename(pocket_path)]*N)350 351 sdf_names += [os.path.basename(ligands_path)]*len(confidence)352 353 assert len(confidence_list)==len(confidence_names)==len(sdf_names)==len(pocket_path_list)354 """ add a save command by caoduanhua to save the last state of ligand """355 ########################################################################356 if args.save_docking_result:357 """"if you use multiple molecule parallel inference, you should re_order the confidence one by one"""358 # add a parm to control the number of save ligand pose359 head_threshold = 0360 tail_threshold = N361 confidence_tmp = confidence[head_threshold:tail_threshold]362 re_order = np.argsort(confidence_tmp)[::-1]363 if args.inference_mode=='evaluate':364 true_mol = remove_all_hs(read_abs_file_mol(ref_ligand))365 for _ in range(len(orig_complex_graph_list)):366 for rank, batch_idx in enumerate(re_order[:args.save_docking_result_number]):367 true_idx = head_threshold + batch_idx368 mol_pred = copy.deepcopy(data_list[true_idx]['mol'])369 370 pos = data_list[true_idx]['ligand'].pos.cpu().numpy() + orig_complex_graph_list[_].original_center.cpu().numpy()371 372 if score_model_args.remove_hs: mol_pred = remove_all_hs(mol_pred)373 374 375 if args.inference_mode=='evaluate':376 try:377 rmsd = get_symmetry_rmsd(true_mol, true_mol.GetConformers()[0].GetPositions(), [pos])[0]378 except Exception as e:379 logger.warning("Using non corrected RMSD because of the error:{}", e)380 381 rmsd = np.sqrt(((true_mol.GetConformers()[0].GetPositions() - pos) ** 2).sum(axis=-1).mean(axis=0))382 result_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_rmsd_{rmsd}_confidence_{confidence_tmp[batch_idx]}.sdf'383 else:384 result_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_confidence_{confidence_tmp[batch_idx]}.sdf'385 386 write_mol_with_coords(mol_pred, pos, os.path.join(write_dir, result_filename))387 388 if args.save_visualisation:389 write_dir_vis = f'{args.out_dir}/SurfDock_docking_result/{data_list[true_idx]["name"]}'390 os.makedirs(write_dir, exist_ok=True)391 if args.inference_mode=='evaluate':392 vis_filename =f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_rmsd_{rmsd}_confidence_{confidence_tmp[batch_idx]}.pdb'393 else:394 vis_filename = f'{data_list[true_idx]["name"]}_sample_idx_{batch_idx}_rank_{rank + 1}_confidence_{confidence_tmp[batch_idx]}.pdb'395 try:396 visualization_list[batch_idx].write(397 f'{write_dir_vis}/{vis_filename}')398 except:399 continue400 head_threshold += N401 tail_threshold += N402 403 if _ < len(orig_complex_graph_list) - 1:404 405 confidence_tmp = confidence[head_threshold:tail_threshold]406 re_order = np.argsort(confidence_tmp)[::-1]407 except Exception as e:408 # if isinstance(e,RecursionError) or 'out of memory' in str(e):409 data_list = None410 referrers = gc.get_referrers(data_list)411 for ref in referrers:412 ref=None413 gc.collect()414 torch.cuda.empty_cache()415 data_list = [copy.deepcopy(temp_graph) for temp_graph in orig_complex_graph_list for _ in range(N)]416 logger.error("Failed on :{}, error of :{}", orig_complex_graph["name"], e)417 failures += 1418 sample_count_failed +=1419 if sample_count_failed > 5:420 logger.error(" Skip by five times Failed on :{}, error of :{}", orig_complex_graph["name"], e)421 success = 1422 else:423 success = 0424 425 except Exception as e:426 if 'out of memory' in str(e):427 logger.critical('| WARNING: ran out of memory, skipping batch')428 orig_complex_graph_list,orig_complex_graph,data_list=None,None,None429 referrers = gc.get_referrers(data_list)430 for ref in referrers:431 ref=None432 gc.collect()433 torch.cuda.empty_cache()434 logger.error('Some error failed for sampling: idx in batch : {}, ligand_path: {},error of :{} ',idx,ligands_path,e)435 436 continue437 # if args.inference_mode=='evaluate':438 esm_embeddings,test_dataset,test_loader,orig_complex_graph_list,orig_complex_graph,data_list=None,None,None,None,None,None439 gc.collect()440 torch.cuda.empty_cache()441 except Exception as e:442 logger.error('Some error failed for graph data. ligand_path: {},error of :{}',ligands_path,e)443 esm_embeddings,test_dataset,test_loader,orig_complex_graph_list,orig_complex_graph,data_list=None,None,None,None,None,None444 referrers = gc.get_referrers(data_list)445 for ref in referrers:446 ref=None447 gc.collect()448 torch.cuda.empty_cache()449 continue450 logger.info('Protein {} used time: {}',os.path.splitext(os.path.basename(pocket_path))[0],time.time() - in_loop_start_time)451 accelerator.wait_for_everyone()452 docking_time = time.time() - start_time453 if accelerator.is_local_main_process:454 logger.info('Docking time used for one moleculer: {}',docking_time/ all_molecules)455 logger.info('Docking time used: {}', docking_time)456 logger.info('Sampling conformers number: {}',all_molecules*args.samples_per_complex)457 logger.info('Output conformers number: {}', all_molecules*args.save_docking_result_number)458 logger.info('Docking output molecule number: {}',all_molecules)459 # logger.info('Docking time used for one moleculer: {}',docking_time/all_molecules)460 461 result = pd.DataFrame({'sdf_name':sdf_names,'confidence':confidence_list,'confidence_name':confidence_names,'pocket_path':pocket_path_list})462 csv_flag = os.path.basename(args.data_csv).split('.')[0]463 result.to_csv(f'{args.out_dir}/{csv_flag}_head_{str(args.head_index)}_tail_{str(args.tail_index)}_confidence_on_device_{device}.csv',index=False)464 465 if accelerator.is_local_main_process:466 if args.wandb:467 wandb.finish()468if __name__ == '__main__':469 from accelerate import Accelerator470 from accelerate.utils import DistributedDataParallelKwargs471 kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)472 accelerator = Accelerator(kwargs_handlers=[kwargs])473 from accelerate.utils import set_seed474 device = accelerator.device475 set_seed(args.random_seed)476 from functools import partial477 478 accelerator.print(f'device {str(accelerator.device)} is used!')479 main_function()480 