from __future__ import annotations import json import os import sys from pathlib import Path os.environ.setdefault("TF_CPP_MIN_LOG_LEVEL", "2") os.environ.setdefault("TF_ENABLE_ONEDNN_OPTS", "0") import numpy as np ROOT = Path(__file__).resolve().parent sys.path.insert(0, str(ROOT / "src")) from forge_lab.models.numpy_runtime import NumpyDriftRuntime # noqa: E402 ARTIFACTS = ROOT def sample_windows() -> np.ndarray: from export_models import synthetic_windows values, _ = synthetic_windows(count=64, seed=91) return values def run() -> dict: import onnxruntime as ort import tensorflow as tf import torch from ai_edge_litert.interpreter import Interpreter from forge_lab.models.pytorch_model import TinyDriftNet windows = sample_windows() numpy_output, _ = NumpyDriftRuntime().predict_batch(windows) checkpoint = torch.load( ARTIFACTS / "tiny_drift_pytorch.pt", map_location="cpu", weights_only=True ) pytorch_model = TinyDriftNet().eval() pytorch_model.load_state_dict(checkpoint["state_dict"]) with torch.no_grad(): pytorch_output = pytorch_model(torch.from_numpy(windows)).numpy() session = ort.InferenceSession( str(ARTIFACTS / "tiny_drift.onnx"), providers=["CPUExecutionProvider"] ) onnx_output = session.run(["probability"], {"telemetry": windows})[0] saved_model = tf.saved_model.load(str(ARTIFACTS / "tensorflow_saved_model")) tensorflow_output = saved_model.signatures["serving_default"](telemetry=tf.constant(windows))[ "probability" ].numpy() interpreter = Interpreter(model_path=str(ARTIFACTS / "tiny_drift.tflite")) input_details = interpreter.get_input_details()[0] output_details = interpreter.get_output_details() interpreter.resize_tensor_input(input_details["index"], windows.shape, strict=False) interpreter.allocate_tensors() interpreter.set_tensor(input_details["index"], windows) interpreter.invoke() probability_detail = next(item for item in output_details if len(item["shape_signature"]) == 1) litert_output = interpreter.get_tensor(probability_detail["index"]) litert_output = litert_output.reshape(-1) runtimes = { "pytorch": pytorch_output, "onnxruntime": onnx_output, "tensorflow": tensorflow_output, "litert": litert_output, } report = {} for name, output in runtimes.items(): max_abs_error = float(np.max(np.abs(numpy_output - output))) agreement = float(np.mean((numpy_output >= 0.5) == (output >= 0.5))) report[name] = { "max_abs_error_vs_numpy": max_abs_error, "classification_agreement": agreement, "passed": max_abs_error <= 1e-4 and agreement == 1.0, } report["metadata"] = { "windows": len(windows), "threshold": 0.5, "max_abs_error_gate": 1e-4, "reference": "numpy", } report["all_passed"] = all(report[name]["passed"] for name in runtimes) return report if __name__ == "__main__": result = run() (ARTIFACTS / "conformance.json").write_text(json.dumps(result, indent=2) + "\n") print(json.dumps(result, indent=2)) raise SystemExit(0 if result["all_passed"] else 1)