NABP-LSTM-Att / scripts /train_one_epoch.py
wuxing0105's picture
Upload folder using huggingface_hub
bf928ee verified
Raw
History Blame Contribute Delete
9.42 kB
#!/usr/bin/env python3
"""Train the full official training set for one epoch on a single DCU."""
import argparse
import json
import pickle
import random
import time
from pathlib import Path
from _bootstrap import DATA_DIR, OUTPUT_DIR
import numpy as np
import tensorflow as tf
import tf_keras
from sklearn.metrics import average_precision_score, roc_auc_score
DEFAULT_DATA_DIR = DATA_DIR / "features" / "cdr_kmer3_ag_kmer1"
DEFAULT_OUTPUT_DIR = OUTPUT_DIR / "one_epoch"
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--data-dir", type=Path, default=DEFAULT_DATA_DIR)
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--seed", type=int, default=123)
return parser.parse_args()
def load_split(data_dir, suffix):
# Register features.InputFeatures before loading the official pickles.
import model # noqa: F401
with (data_dir / f"cdr_features_{suffix}.pickle").open("rb") as reader:
cdr_features = pickle.load(reader)
with (data_dir / f"ag_features_{suffix}.pickle").open("rb") as reader:
ag_features = pickle.load(reader)
if not cdr_features:
raise ValueError(f"The {suffix} split is empty")
if len(cdr_features) != len(ag_features):
raise ValueError(
f"Unpaired {suffix} features: CDR={len(cdr_features)}, "
f"AG={len(ag_features)}"
)
cdr_labels = np.asarray(
[[item.label_id] for item in cdr_features], dtype=np.float32
)
ag_labels = np.asarray(
[[item.label_id] for item in ag_features], dtype=np.float32
)
if not np.array_equal(cdr_labels, ag_labels):
mismatch_count = int(np.count_nonzero(cdr_labels != ag_labels))
raise ValueError(f"CDR/AG labels differ for {mismatch_count} {suffix} samples")
if set(np.unique(cdr_labels).tolist()) != {0.0, 1.0}:
raise ValueError(f"Expected binary labels 0/1, got {np.unique(cdr_labels)}")
inputs = [
np.asarray([item.input_ids for item in cdr_features], dtype=np.int32),
np.asarray([item.cdr_number_ids for item in cdr_features], dtype=np.int32),
np.asarray([item.input_ids for item in ag_features], dtype=np.int32),
]
sample_count = len(cdr_features)
expected_shapes = ((sample_count, 24), (sample_count, 24), (sample_count, 2371))
actual_shapes = tuple(array.shape for array in inputs)
if actual_shapes != expected_shapes:
raise ValueError(f"Unexpected {suffix} shapes: {actual_shapes}, expected {expected_shapes}")
return inputs, cdr_labels
def main():
args = parse_args()
if args.batch_size < 2:
raise ValueError(f"--batch-size must be at least 2, got {args.batch_size}")
if not args.data_dir.is_dir():
raise FileNotFoundError(f"Data directory not found: {args.data_dir}")
# Never reuse an existing directory: this prevents accidental checkpoint overwrite.
args.output_dir.mkdir(parents=True, exist_ok=False)
random.seed(args.seed)
np.random.seed(args.seed)
tf.random.set_seed(args.seed)
gpus = tf.config.list_physical_devices("GPU")
if not gpus:
raise RuntimeError(
"TensorFlow cannot see a DCU. Set HIP_VISIBLE_DEVICES and "
"CUDA_VISIBLE_DEVICES before starting Python."
)
for gpu in gpus:
tf.config.experimental.set_memory_growth(gpu, True)
train_inputs, train_labels = load_split(args.data_dir, "tr")
val_inputs, val_labels = load_split(args.data_dir, "val")
# Import after configuring the accelerator to avoid initializing it early.
from model import get_model
with tf.device("/GPU:0"):
model = get_model()
initial_weights = [
variable.numpy().copy() for variable in model.trainable_variables
]
start_time = time.perf_counter()
history = model.fit(
train_inputs,
train_labels,
validation_data=(val_inputs, val_labels),
epochs=1,
batch_size=args.batch_size,
shuffle=True,
verbose=1,
)
training_seconds = time.perf_counter() - start_time
val_predictions = model.predict(
val_inputs,
batch_size=args.batch_size,
verbose=1,
).reshape(-1)
if not np.all(np.isfinite(val_predictions)):
raise FloatingPointError("Validation predictions contain NaN or Inf")
if np.any((val_predictions < 0.0) | (val_predictions > 1.0)):
raise ValueError("Validation predictions are outside [0, 1]")
training_loss = float(history.history["loss"][-1])
validation_loss = float(history.history["val_loss"][-1])
if not np.isfinite(training_loss) or not np.isfinite(validation_loss):
raise FloatingPointError(
f"Non-finite loss: train={training_loss}, val={validation_loss}"
)
changed_count = 0
max_weight_change = 0.0
for before, variable in zip(initial_weights, model.trainable_variables):
after = variable.numpy()
if not np.all(np.isfinite(after)):
raise FloatingPointError(f"Trainable variable contains NaN/Inf: {variable.name}")
change = float(np.max(np.abs(after - before)))
max_weight_change = max(max_weight_change, change)
if change > 0.0:
changed_count += 1
labels_flat = val_labels.reshape(-1).astype(np.int32)
val_auroc = float(roc_auc_score(labels_flat, val_predictions))
val_aupr = float(average_precision_score(labels_flat, val_predictions))
dcu_variable_count = sum("GPU:0" in variable.device for variable in model.variables)
if changed_count == 0:
raise AssertionError("No trainable weights changed during the epoch")
if dcu_variable_count == 0:
raise RuntimeError("No model variables were placed on GPU:0")
checkpoint_path = args.output_dir / "model_after_one_epoch.weights.h5"
predictions_path = args.output_dir / "validation_predictions.npz"
metrics_path = args.output_dir / "metrics.json"
model.save_weights(str(checkpoint_path))
np.savez_compressed(
predictions_path,
labels=labels_flat,
predictions=val_predictions,
)
# Verify that the newly written checkpoint can be loaded and reproduces output.
with tf.device("/GPU:0"):
reloaded_model = get_model()
reloaded_model.load_weights(str(checkpoint_path))
reload_count = min(args.batch_size, len(labels_flat))
reloaded_predictions = reloaded_model(
[array[:reload_count] for array in val_inputs],
training=False,
).numpy().reshape(-1)
reload_max_error = float(
np.max(np.abs(reloaded_predictions - val_predictions[:reload_count]))
)
if reload_max_error > 1e-6:
raise AssertionError(
f"Reloaded checkpoint max error {reload_max_error:.10g} exceeds 1e-6"
)
metrics = {
"tensorflow_version": tf.__version__,
"tf_keras_version": tf_keras.__version__,
"seed": args.seed,
"batch_size": args.batch_size,
"epochs": 1,
"initialization": "random",
"train_samples": int(train_labels.size),
"validation_samples": int(val_labels.size),
"training_loss": training_loss,
"validation_loss": validation_loss,
"validation_auroc": val_auroc,
"validation_aupr": val_aupr,
"training_seconds": training_seconds,
"throughput_samples_per_second": float(train_labels.size / training_seconds),
"trainable_tensors_changed": changed_count,
"trainable_tensor_count": len(initial_weights),
"max_absolute_weight_change": max_weight_change,
"variables_on_gpu0": dcu_variable_count,
"model_variable_count": len(model.variables),
"checkpoint_reload_max_error": reload_max_error,
}
metrics_path.write_text(
json.dumps(metrics, indent=2, ensure_ascii=False) + "\n",
encoding="utf-8",
)
print("\nFull training-set one-epoch validation")
print(f"TensorFlow: {tf.__version__}")
print(f"tf_keras: {tf_keras.__version__}")
print(f"DCU: {gpus[0]}")
print("Initialization: random (official Model99.h5 was not loaded)")
print(f"Train/validation samples: {train_labels.size}/{val_labels.size}")
print(f"Batch size/epochs: {args.batch_size}/1")
print(f"Training loss: {training_loss:.10g}")
print(f"Validation loss: {validation_loss:.10g}")
print(f"Validation AUROC: {val_auroc:.10f}")
print(f"Validation AUPR: {val_aupr:.10f}")
print(f"Training seconds: {training_seconds:.3f}")
print(f"Training throughput samples/s: {train_labels.size / training_seconds:.3f}")
print(f"Trainable tensors changed: {changed_count}/{len(initial_weights)}")
print(f"Max absolute weight change: {max_weight_change:.10g}")
print(f"Variables placed on GPU:0: {dcu_variable_count}/{len(model.variables)}")
print(f"Checkpoint reload max error: {reload_max_error:.10g}")
print(f"Checkpoint: {checkpoint_path}")
print(f"Metrics: {metrics_path}")
print(f"Validation predictions: {predictions_path}")
print("Full training-set one epoch: PASS")
print("Official Model99.h5 was not loaded or overwritten.")
if __name__ == "__main__":
main()