Skip to content

Commit 8b758c9

Browse files
Merge pull request #97 from IBM/apml-2
🛠️ Make apml.build_model, renunable
2 parents f4c56ab + 8f6a355 commit 8b758c9

3 files changed

Lines changed: 803 additions & 81 deletions

File tree

autopeptideml/apml.py

Lines changed: 42 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,10 @@ class AutoPeptideML:
4444
"""
4545
df: pd.DataFrame
4646
metadata: dict = {}
47+
parts = None
48+
x: dict = {}
49+
execution: dict = {}
50+
hpo_run: int = 1
4751

4852
def __init__(
4953
self,
@@ -269,6 +273,8 @@ def build_models(
269273
"Please try: `min`",
270274
"`good` strategy will be implemented in future releases."
271275
)
276+
if osp.isdir(osp.join(self.outputdir, 'ensemble')):
277+
shutil.rmtree(osp.join(self.outputdir, 'ensemble'))
272278
self._partitioning(
273279
split_strategy=split_strategy,
274280
hestia_generator=hestia_generator,
@@ -355,8 +361,9 @@ def _hpo(
355361
random_state=random_state,
356362
n_jobs=n_jobs,
357363
db_file=osp.join(self.meta_dir, 'database.sql'),
358-
study_name='apml-1'
364+
study_name=f'apml-{self.hpo_run}'
359365
)
366+
self.hpo_run += 1
360367
end = time.time()
361368
self.metadata['status'] = 'trained'
362369
self.metadata['trainer-metadata'].update({
@@ -365,7 +372,8 @@ def _hpo(
365372
'metric': metric,
366373
'n-trials': n_trials,
367374
'patience': n_trials // 5,
368-
'best-run': int(self.trainer.best_run)
375+
'best-run': int(self.trainer.best_run),
376+
'hpo-run': self.hpo_run
369377
})
370378
input_trial = {rep: self.x[rep][:1]
371379
for rep in self.trainer.best_model.reps}
@@ -383,7 +391,15 @@ def _representing(
383391
n_jobs: int,
384392
verbose: bool
385393
):
386-
self.x, execution = {}, {}
394+
all_available = True
395+
for rep in reps:
396+
if rep not in self.x:
397+
all_available = False
398+
break
399+
400+
if all_available:
401+
return
402+
reps = [r for r in reps if r not in self.x]
387403
prot, mol = False, False
388404

389405
for rep in reps:
@@ -404,11 +420,11 @@ def _representing(
404420

405421
if isinstance(reps, dict):
406422
for name, repengine in reps.items():
407-
execution[name] = {'start': time.time()}
423+
self.execution[name] = {'start': time.time()}
408424
self.x[name] = repengine.compute_reps(
409425
self.df[self.sequence_field], verbose=verbose, batch_size=16
410426
)
411-
execution[name]['end'] = time.time()
427+
self.execution[name]['end'] = time.time()
412428
reps = list(reps.keys())
413429
else:
414430
for rep in reps:
@@ -417,7 +433,7 @@ def _representing(
417433

418434
if rep in PLMs or rep in CLMs:
419435
from autopeptideml.reps.lms import RepEngineLM
420-
execution[rep] = {'start': time.time()}
436+
self.execution[rep] = {'start': time.time()}
421437

422438
repengine = RepEngineLM(rep, average_pooling=True,
423439
fp16=True)
@@ -429,59 +445,64 @@ def _representing(
429445
rep = f'{rep}-8-2048'
430446
elif len(rep.split('-')) == 2:
431447
rep = f"{rep.split('-')[0]}-{rep.split('-')[1]}-2048"
432-
execution[rep] = {'start': time.time()}
448+
self.execution[rep] = {'start': time.time()}
433449
repengine = RepEngineFP(
434450
rep=rep.split('-')[0],
435451
radius=int(rep.split('-')[1]),
436452
nbits=int(rep.split('-')[2])
437453
)
438454
elif rep == 'one-hot':
439455
from autopeptideml.reps.seq_based import RepEngineOnehot
440-
execution[rep] = {'start': time.time()}
456+
self.execution[rep] = {'start': time.time()}
441457

442458
repengine = RepEngineOnehot(max_length=50)
443459

444460
batch_size = 128 if repengine.get_num_params() < 2e7 else 16
445461
if rep in PLMs or rep == 'one-hot':
446462
if 'to-sequences' in self.metadata['pipeline-1']['name']:
447-
execution[rep] = {'start': time.time()}
463+
self.execution[rep] = {'start': time.time()}
448464

449465
self.x[rep] = repengine.compute_reps(
450466
self.df[f'{self.sequence_field}'], verbose=verbose,
451467
batch_size=batch_size
452468
)
453469
else:
454-
execution[rep] = {'start': time.time()}
470+
self.execution[rep] = {'start': time.time()}
455471

456472
self.x[rep] = repengine.compute_reps(
457473
self.df[f'{self.sequence_field}-2'], verbose=verbose,
458474
batch_size=batch_size
459475
)
460476
elif rep in CLMs or rep.split('-')[0] in FPs:
461477
if 'to-smiles' in self.metadata['pipeline-1']['name']:
462-
execution[rep] = {'start': time.time()}
478+
self.execution[rep] = {'start': time.time()}
463479

464480
self.x[rep] = repengine.compute_reps(
465481
self.df[f'{self.sequence_field}'], verbose=verbose,
466482
batch_size=batch_size
467483
)
468484
else:
469-
execution[rep] = {'start': time.time()}
485+
self.execution[rep] = {'start': time.time()}
470486

471487
self.x[rep] = repengine.compute_reps(
472-
self.df[f'{self.sequence_field}-2'], verbose=verbose,
488+
self.df[f'{self.sequence_field}-2'],
489+
verbose=verbose,
473490
batch_size=batch_size
474491
)
475-
execution[rep]['end'] = time.time()
492+
self.execution[rep]['end'] = time.time()
476493

477-
self.x = {rep: np.array(value) for rep, value in self.x.items()}
494+
self.x.update({rep: np.array(value) for rep, value in self.x.items()})
478495
path = osp.join(self.meta_dir, 'reps.pckl')
479496
pickle.dump(self.x, open(path, 'wb'))
480497
self.metadata['status'] = 'represented'
481-
self.metadata['reps-metadata'] = {'reps': list(self.x.keys())}
498+
if 'reps-metadata' in self.metadata:
499+
self.metadata['reps-metadata'].update({'reps': list(self.x.keys())})
500+
else:
501+
self.metadata['reps-metadata'] = {'reps': list(self.x.keys())}
502+
482503
self.metadata['reps-metadata'].update({
483504
f'{rep}-execution-time':
484-
execution[rep]['end'] - execution[rep]['start']
505+
self.execution[rep]['end'] - self.execution[rep]['start']
485506
for rep in self.x.keys()})
486507
self.save_metadata()
487508

@@ -501,7 +522,10 @@ def _partitioning(
501522
part_path = osp.join(self.outputdir, 'parts.pckl')
502523
parts = {k: v for k, v in partitions.items()}
503524
pickle.dump(parts, open(part_path, 'wb'))
504-
self.parts = parts
525+
self.parts = partitions
526+
return
527+
if self.parts is not None:
528+
return
505529

506530
SPLIT_STRATEGIES = ['random', 'min', 'good', None]
507531
self.metadata['status'] = 'partitioning'

0 commit comments

Comments
 (0)