NABP-LSTM-Att / scripts /verify_minimal_inference.py
wuxing0105's picture
Upload folder using huggingface_hub
bf928ee verified
Raw
History Blame Contribute Delete
4.72 kB
#!/usr/bin/env python3
"""Run one or more official test samples on CPU and DCU and compare outputs."""
import argparse
import pickle
from pathlib import Path
from _bootstrap import DATA_DIR, WEIGHT_DIR
import numpy as np
import tensorflow as tf
DEFAULT_DATA_DIR = DATA_DIR / "features" / "cdr_kmer3_ag_kmer1"
DEFAULT_WEIGHT_PATH = WEIGHT_DIR / "cdr_kmer3_ag_kmer1" / "Model99.h5"
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--data-dir", type=Path, default=DEFAULT_DATA_DIR)
parser.add_argument("--weights", type=Path, default=DEFAULT_WEIGHT_PATH)
parser.add_argument("--samples", type=int, default=1)
parser.add_argument("--atol", type=float, default=1e-3)
return parser.parse_args()
def load_inputs(data_dir, sample_count):
# Importing features registers features.InputFeatures for pickle loading.
import model # noqa: F401
with (data_dir / "cdr_features_te.pickle").open("rb") as reader:
cdr_features = pickle.load(reader)
with (data_dir / "ag_features_te.pickle").open("rb") as reader:
ag_features = pickle.load(reader)
if len(cdr_features) != len(ag_features):
raise ValueError(
f"Unpaired test features: CDR={len(cdr_features)}, AG={len(ag_features)}"
)
if sample_count < 1 or sample_count > len(cdr_features):
raise ValueError(
f"--samples must be between 1 and {len(cdr_features)}, got {sample_count}"
)
selected_cdr = cdr_features[:sample_count]
selected_ag = ag_features[:sample_count]
inputs = [
np.asarray([item.input_ids for item in selected_cdr], dtype=np.int32),
np.asarray([item.cdr_number_ids for item in selected_cdr], dtype=np.int32),
np.asarray([item.input_ids for item in selected_ag], dtype=np.int32),
]
labels = np.asarray([item.label_id for item in selected_cdr], dtype=np.float32)
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 input shapes: {actual_shapes}, expected {expected_shapes}")
return inputs, labels
def predict_on_device(device, inputs, weight_path):
from model import get_model
with tf.device(device):
model = get_model()
model.load_weights(str(weight_path))
prediction_tensor = model(inputs, training=False)
predictions = prediction_tensor.numpy().reshape(-1)
if predictions.shape != (inputs[0].shape[0],):
raise ValueError(f"Unexpected prediction shape: {predictions.shape}")
if not np.all(np.isfinite(predictions)):
raise FloatingPointError(f"{device} predictions contain NaN or Inf")
if np.any((predictions < 0.0) | (predictions > 1.0)):
raise ValueError(f"{device} predictions are outside [0, 1]")
return predictions, prediction_tensor.device
def main():
args = parse_args()
if not args.data_dir.is_dir():
raise FileNotFoundError(f"Data directory not found: {args.data_dir}")
if not args.weights.is_file():
raise FileNotFoundError(f"Weight file not found: {args.weights}")
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)
inputs, labels = load_inputs(args.data_dir, args.samples)
cpu_predictions, cpu_device = predict_on_device("/CPU:0", inputs, args.weights)
dcu_predictions, dcu_device = predict_on_device("/GPU:0", inputs, args.weights)
absolute_error = np.abs(cpu_predictions - dcu_predictions)
max_error = float(np.max(absolute_error))
mean_error = float(np.mean(absolute_error))
print(f"TensorFlow: {tf.__version__}")
print(f"Samples: {args.samples}")
print(f"CPU tensor device: {cpu_device}")
print(f"DCU tensor device: {dcu_device}")
print(f"Labels: {labels.tolist()}")
print(f"CPU predictions: {cpu_predictions.tolist()}")
print(f"DCU predictions: {dcu_predictions.tolist()}")
print(f"Max absolute error: {max_error:.10g}")
print(f"Mean absolute error: {mean_error:.10g}")
if "GPU:0" not in dcu_device:
raise RuntimeError(f"DCU inference was not placed on GPU:0: {dcu_device}")
if max_error > args.atol:
raise AssertionError(
f"CPU/DCU max absolute error {max_error:.10g} exceeds --atol {args.atol}"
)
print("Minimal CPU/DCU inference: PASS")
if __name__ == "__main__":
main()