MS2-SMILES-AlignNet / config.py
monaaaaaa's picture
Upload 26 files
eeabcff verified
Raw
History Blame Contribute Delete
3.49 kB
import torch, json, math, os
d = {
'debug': True,
'dataset_path': '../data/tongji_data/pos.json',
'train_data_file': '../data/tongji_data/tmp_input_ids_27/',
'train_number_data': 1632069,
'valid_data_file': '../data/tongji_data/test_2016_new_input_id_new.pt',
'fptype': 'morgan',
'valid_ratio': 0.1,
'batch_size': 64,
'lr': 1e-3,
'min_lr': 2e-5,
'weight_decay': 1e-3,
'scheduler_type': 'warmup_cosine',
'warmup_epochs': 1,
'warmup_start_factor': 1.0,
'patience': 2,
'factor': 0.5,
'add_nl': True,
'binary_intn': False,
'max_mz': 2000,
'min_mz': 20,
'energy': 'Energy1',
'epochs': 50,
'bin_size': 0.05,
'ms_embedding_dim': 300,
'ms_feature3_embedding_dim': 135,
'ms3_meta_dim': 50,
'projection_dim': 256,
'ms_projection_layers': 1,
'mol_embedding_dim': 2048,
'mol_projection_layers': 1,
'tsfm_in_ms': True,
'tsfm_in_mol': False,
'tsfm_layers': 6,
'tsfm_heads': 8,
'lstm_layers': 2,
'lstm_in_ms': False,
'lstm_in_mol': False,
'dropout': 0.1,
'nmodels': 1,
'mol_encoder': 'gnn+fp', # fp, gnn or gnn+fp
'molgnn_n_filters_list': [256, 256, 256],
'molgnn_nhead': 4,
'molgnn_readout_layers': 2,
'seed': 1234,
'dev_name': 'cuda:0',
'keep_best_models_num': 3,
'alpha': 0.5,
'beta': 0.5,
'structure_diag_weight': 0.25,
'fp_dim': 2048,
"experiment_name_type": ['mol_gcn_base', 'mol_gat_base', 'mol_gcn_ms1', 'mol_gat_ms1', 'mol_gcn_ms3', 'mol_gat_ms3',
'mol_gcn_ms1_ms3', 'mol_gat_ms1_ms3',
# loss 8 - 15
'mol_gcn_base_loss', 'mol_gat_base_loss', 'mol_gcn_ms1_loss', 'mol_gat_ms1_loss',
'mol_gcn_ms3_loss', 'mol_gat_ms3_loss', 'mol_gcn_ms1_ms3_loss', 'mol_gat_ms1_ms3_loss'][14]
}
class ConfigDict(dict):
'''
Makes a dictionary behave like an object,with attribute-style access.
'''
def __getattr__(self, name):
try:
return self[name]
except:
raise AttributeError(name)
def __setattr__(self, name, value):
self[name] = value
def save(self, fn, onlyprint=False):
if onlyprint:
print(self)
else:
json.dump(self, open(fn, 'w'), indent=2)
def load_dict(self, dic):
for k, v in dic.items():
self[k] = v
self.calc_ms_embedding_dim()
def load(self, fn):
try:
if type(fn) is dict:
d = fn
elif type(fn) is str:
if os.path.exists(fn):
d = json.load(open(fn, 'r'))
else:
d = json.loads(fn)
self.load_dict(d)
except Exception as e:
print(e)
def calc_ms_embedding_dim(self):
if 'bin_size' in self:
self['ms_embedding_dim'] = math.ceil((self['max_mz'] - self['min_mz']) / self['bin_size'])
if 'ms_embedding_dim' in self and 'add_nl' in self and self['add_nl']:
self['ms_embedding_dim'] += math.ceil((200) / self['bin_size'])
@property
def device(self):
try:
if torch.cuda.is_available():
return torch.device(self['dev_name'])
else:
return torch.device('cpu')
except Exception as e:
return torch.device('cpu')
CFG = ConfigDict()
CFG.load_dict(d)