Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Binary file modified src/geneml/models/geneML_default.keras
Binary file not shown.
152 changes: 132 additions & 20 deletions trainer/train_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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')))

Expand All @@ -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()

###############################################################################
Expand All @@ -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))
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()

Expand All @@ -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))
Expand Down
Loading