diff --git a/src/geneml/models/geneML_default.keras b/src/geneml/models/geneML_default.keras index 13437fd..d95626b 100644 Binary files a/src/geneml/models/geneML_default.keras and b/src/geneml/models/geneML_default.keras differ diff --git a/trainer/train_model.py b/trainer/train_model.py index 3ab133d..19077bf 100755 --- a/trainer/train_model.py +++ b/trainer/train_model.py @@ -79,12 +79,24 @@ def get_args(): default=0.0005, help='Minimum improvement in validation metric to reset patience (default: %(default)s)', ) + parser.add_argument( + '--learning-rate-decay', + type=float, + default=1/3, + help='Factor to decay learning rate by on plateau (default: %(default)s)', + ) parser.add_argument( '--chunks-per-batch', type=int, default=1, help='How many dataset chunks to pool, shuffle, and train on in a single fit call (default: 1)' ) + parser.add_argument( + '--resume-from-checkpoint', + type=str, + default=None, + help='Path to .keras checkpoint file to resume training from (default: None)' + ) args, _ = parser.parse_known_args() return args @@ -134,6 +146,33 @@ def get_config(self): } + @keras.saving.register_keras_serializable() + class DecayOnPlateauSchedule(tensorflow.keras.optimizers.schedules.LearningRateSchedule): + def __init__(self, base_schedule, decay_factor): + self.base_schedule = base_schedule + self.decay_factor = decay_factor + + def __call__(self, step): + return self.base_schedule(step) * self.decay_factor + + def get_config(self): + return { + 'base_schedule': keras.saving.serialize_keras_object(self.base_schedule), + 'decay_factor': float(tensorflow.keras.backend.get_value(self.decay_factor)), + } + + @classmethod + def from_config(cls, config): + base_schedule = keras.saving.deserialize_keras_object(config['base_schedule']) + decay_factor = tensorflow.Variable( + config['decay_factor'], + trainable=False, + dtype=tensorflow.float32, + name='lr_decay_factor', + ) + return cls(base_schedule, decay_factor) + + args = get_args() os.makedirs('Models', exist_ok=True) @@ -227,11 +266,13 @@ def categorical_crossentropy_2d_gene_ml(y_true, y_pred): # Set learning rate schedule with warmup steps_per_epoch = num_idx_train // BATCH_SIZE warmup_steps = WARMUP_EPOCHS * steps_per_epoch - lr_schedule = WarmupSchedule( + base_lr_schedule = WarmupSchedule( start_lr=LEARNING_RATE * 0.01, target_lr=LEARNING_RATE, warmup_steps=warmup_steps, ) + decay_factor = tensorflow.Variable(1.0, trainable=False, dtype=tensorflow.float32, name='lr_decay_factor') + lr_schedule = DecayOnPlateauSchedule(base_lr_schedule, decay_factor) print("Num GPUs Available: ", len(tensorflow.config.list_physical_devices('GPU'))) @@ -243,20 +284,75 @@ def categorical_crossentropy_2d_gene_ml(y_true, y_pred): except Exception: pass - if N_GPUS > 1: - # https://keras.io/guides/distributed_training/ - strategy = tensorflow.distribute.MirroredStrategy() - tee('Number of devices: {}'.format(strategy.num_replicas_in_sync)) - with strategy.scope(): + if args.resume_from_checkpoint: + tee(f"\033[1mLoading model from checkpoint: {args.resume_from_checkpoint}\033[0m") + # Load with custom objects + custom_objects = { + 'categorical_crossentropy_2d_gene_ml': categorical_crossentropy_2d_gene_ml, + 'WarmupSchedule': WarmupSchedule, + 'DecayOnPlateauSchedule': DecayOnPlateauSchedule, + } + model = keras.models.load_model(args.resume_from_checkpoint, custom_objects=custom_objects) + + # Extract the decay_factor Variable directly from the loaded schedule + # The optimizer._learning_rate holds the actual schedule object + loaded_schedule = model.optimizer._learning_rate + if hasattr(loaded_schedule, 'decay_factor'): + # Use the existing Variable from the loaded model + decay_factor = loaded_schedule.decay_factor + current_decay = float(tensorflow.keras.backend.get_value(decay_factor)) + tee(f"Loaded decay_factor from checkpoint: {current_decay}") + else: + # Fallback: reconstruct the schedule with a new decay_factor + tee("Warning: Could not find decay_factor in loaded schedule, reconstructing") + try: + # Try to get the saved decay value from config + lr_config = model.optimizer.learning_rate.get_config() + saved_decay = lr_config.get('decay_factor', 1.0) + except (AttributeError, KeyError, TypeError): + saved_decay = 1.0 + + # Reconstruct the learning rate schedule + base_lr_schedule = WarmupSchedule( + start_lr=LEARNING_RATE * 0.01, + target_lr=LEARNING_RATE, + warmup_steps=warmup_steps, + ) + decay_factor = tensorflow.Variable(saved_decay, trainable=False, dtype=tensorflow.float32, name='lr_decay_factor') + new_lr_schedule = DecayOnPlateauSchedule(base_lr_schedule, decay_factor) + + # Replace the optimizer's learning rate schedule + model.optimizer.learning_rate = new_lr_schedule + tee(f"Reconstructed schedule with decay_factor: {saved_decay}") + + # Get current learning rate + current_lr = model.optimizer.learning_rate + if callable(current_lr): + lr_value = float(tensorflow.keras.backend.get_value(current_lr(model.optimizer.iterations))) + else: + lr_value = float(tensorflow.keras.backend.get_value(current_lr)) + tee(f"Resumed from checkpoint. Current learning rate: {lr_value:.5f}") + + # Extract starting epoch from checkpoint filename (e.g., "ep10" -> start at 11) + match = re.search(r'_ep(\d+)\.keras', args.resume_from_checkpoint) + start_epoch = int(match.group(1)) + 1 if match else 1 + tee(f"\033[1mResuming from epoch {start_epoch}\033[0m") + else: + if N_GPUS > 1: + # https://keras.io/guides/distributed_training/ + strategy = tensorflow.distribute.MirroredStrategy() + tee('Number of devices: {}'.format(strategy.num_replicas_in_sync)) + with strategy.scope(): + model = GeneML(L, W, AR, num_classes) + model.compile(loss=loss, + optimizer=keras.optimizers.Adam(learning_rate=lr_schedule, + weight_decay=WEIGHT_DECAY)) + else: model = GeneML(L, W, AR, num_classes) model.compile(loss=loss, optimizer=keras.optimizers.Adam(learning_rate=lr_schedule, weight_decay=WEIGHT_DECAY)) - else: - model = GeneML(L, W, AR, num_classes) - model.compile(loss=loss, - optimizer=keras.optimizers.Adam(learning_rate=lr_schedule, - weight_decay=WEIGHT_DECAY)) + start_epoch = 1 # model.summary() ############################################################################### @@ -279,6 +375,7 @@ def categorical_crossentropy_2d_gene_ml(y_true, y_pred): best_val_score = -1.0 best_epoch = 0 patience_counter = 0 + decay_counter = 0 # Evaluation batch size can be smaller than training to reduce peak GPU mem EVAL_BATCH_SIZE = max(1, min(BATCH_SIZE, args.eval_batch_size * N_GPUS)) @@ -338,7 +435,7 @@ def print_performance_metrics(indices, max_eval): return acceptor_score, donor_score - for epoch_num in range(1, args.num_epochs + 1): + for epoch_num in range(start_epoch, args.num_epochs + 1): # Shuffle indices for this epoch (no replacement) # Use epoch-based seed for deterministic but different permutation per epoch np.random.seed(SEED + epoch_num) @@ -397,9 +494,15 @@ def print_performance_metrics(indices, max_eval): patience_counter += 1 tee(f"No improvement for {patience_counter} evaluation(s). Best: {best_val_score:.4f} (epoch {best_epoch})") - from tensorflow.python.keras import backend - K = backend - tee("Learning rate: %.5f" % (K.get_value(model.optimizer.learning_rate))) + def get_current_lr(): + lr_obj = model.optimizer.learning_rate + if hasattr(lr_obj, "__call__") and not hasattr(lr_obj, "dtype"): + lr_val = lr_obj(model.optimizer.iterations) + else: + lr_val = lr_obj + return float(tensorflow.keras.backend.get_value(lr_val)) + + tee("Learning rate: %.5f" % (get_current_lr())) tee("--- %s seconds ---" % (time.time() - start_time)) start_time = time.time() @@ -423,11 +526,20 @@ def print_performance_metrics(indices, max_eval): # Early stopping: halt if patience exceeded if args.early_stopping_patience > 0 and patience_counter >= args.early_stopping_patience: - tee(f"\n\033[93mEarly stopping triggered after {patience_counter} evaluations without improvement.\033[0m") - tee(f"\033[92mBest epoch: {best_epoch} with validation score: {best_val_score:.4f}\033[0m") - tee(f"\033[92mRecommended checkpoint: GeneML{args.context_length}_c{args.dataset_name}_ep{best_epoch}.keras\033[0m") - h5f.close() - sys.exit(0) + # allow learning rate to decay up to 2 times before stopping + if decay_counter < 2: + old_lr = get_current_lr() + decay_factor.assign(decay_factor * args.learning_rate_decay) + new_lr = old_lr * args.learning_rate_decay + tee(f"\n\033[93mLearning rate decayed from {old_lr:.5f} to {new_lr:.5f} after {patience_counter} evaluations without improvement.\033[0m") + patience_counter = 0 + decay_counter += 1 + else: + tee(f"\n\033[93mEarly stopping triggered after {patience_counter} evaluations without improvement.\033[0m") + tee(f"\033[92mBest epoch: {best_epoch} with validation score: {best_val_score:.4f}\033[0m") + tee(f"\033[92mRecommended checkpoint: GeneML{args.context_length}_c{args.dataset_name}_ep{best_epoch}.keras\033[0m") + h5f.close() + sys.exit(0) else: tee(f"Skipping evaluation this epoch (eval_every={args.eval_every})") tee("--- %s seconds ---" % (time.time() - start_time))