#!/usr/bin/env python3 import os import pickle import sys from collections import defaultdict from typing import Any import numpy as np from tabulate import tabulate from iqpilot.system.hardware import PC from iqpilot.tools.lib.logreader import get_url from iqpilot.selfdrive.test.process_replay.compare_logs import compare_logs, format_diff from iqpilot.selfdrive.test.process_replay.process_replay import get_process_config, replay_process from iqpilot.tools.lib.framereader import FrameReader from iqpilot.tools.lib.logreader import LogReader, save_log TEST_ROUTE = "8494c69d3c710e81|000001d4--2648a9a404" SEGMENT = 4 START_FRAME = 0 END_FRAME = 60 SEND_EXTRA_INPUTS = bool(int(os.getenv("SEND_EXTRA_INPUTS", "0"))) MODEL_REPLAY_BUCKET="model_replay_master" EXEC_TIMINGS = [ # model, instant max, average max ("modelV2", 0.035, 0.025), ("driverStateV2", 0.02, 0.015), ] def get_log_fn(test_route, ref="master"): return f"{test_route}_model_tici_{ref}.zst" def get_model_replay_url(filename): return f"https://raw.githubusercontent.com/commaai/ci-artifacts/refs/heads/{MODEL_REPLAY_BUCKET}/{filename}" def trim_logs(logs, start_frame, end_frame, frs_types, include_all_types): all_msgs = [] cam_state_counts = defaultdict(int) for msg in sorted(logs, key=lambda m: m.logMonoTime): if msg.which() in frs_types: cam_state_counts[msg.which()] += 1 if any(cam_state_counts[state] >= start_frame for state in frs_types): all_msgs.append(msg) if all(cam_state_counts[state] == end_frame for state in frs_types): break if len(include_all_types) != 0: other_msgs = [m for m in logs if m.which() in include_all_types] all_msgs.extend(other_msgs) return all_msgs def model_replay(lr, frs): # modeld is using frame pairs modeld_logs = trim_logs(lr, START_FRAME, END_FRAME, {"roadCameraState", "wideRoadCameraState"}, {"roadEncodeIdx", "wideRoadEncodeIdx", "carParams", "carState", "carControl", "can"}) dmodeld_logs = trim_logs(lr, START_FRAME, END_FRAME, {"driverCameraState"}, {"driverEncodeIdx", "carParams", "can"}) if not SEND_EXTRA_INPUTS: modeld_logs = [msg for msg in modeld_logs if msg.which() != 'extrinsicsCalibration'] dmodeld_logs = [msg for msg in dmodeld_logs if msg.which() != 'extrinsicsCalibration'] # initial setup for s in ('extrinsicsCalibration', 'deviceState'): msg = next(msg for msg in lr if msg.which() == s).as_builder() msg.logMonoTime = lr[0].logMonoTime modeld_logs.insert(1, msg.as_reader()) dmodeld_logs.insert(1, msg.as_reader()) modeld = get_process_config("modeld") dmonitoringmodeld = get_process_config("dmonitoringmodeld") modeld_msgs = replay_process(modeld, modeld_logs, frs) dmonitoringmodeld_msgs = replay_process(dmonitoringmodeld, dmodeld_logs, frs) msgs = modeld_msgs + dmonitoringmodeld_msgs header = ['model', 'max instant', 'max instant allowed', 'average', 'max average allowed', 'test result'] rows = [] timings_ok = True for (s, instant_max, avg_max) in EXEC_TIMINGS: ts = [getattr(m, s).modelExecutionTime for m in msgs if m.which() == s] # TODO some init can happen in first iteration ts = ts[1:] errors = [] if np.max(ts) > instant_max: errors.append("❌ FAILED MAX TIMING CHECK ❌") if np.mean(ts) > avg_max: errors.append("❌ FAILED AVG TIMING CHECK ❌") timings_ok = not errors and timings_ok rows.append([s, np.max(ts), instant_max, np.mean(ts), avg_max, "\n".join(errors) or "✅"]) print("------------------------------------------------") print("----------------- Model Timing -----------------") print("------------------------------------------------") print(tabulate(rows, header, tablefmt="simple_grid", stralign="center", numalign="center", floatfmt=".4f")) assert timings_ok or PC return msgs def get_frames(): regen_cache = "--regen-cache" in sys.argv frames_cache = '/tmp/model_replay_cache' if PC else '/data/model_replay_cache' os.makedirs(frames_cache, exist_ok=True) cache_name = f'{frames_cache}/{TEST_ROUTE}_{SEGMENT}_{START_FRAME}_{END_FRAME}.pkl' if os.path.isfile(cache_name) and not regen_cache: try: print(f"Loading frames from cache {cache_name}") return pickle.load(open(cache_name, "rb")) except Exception as e: print(f"Failed to load frames from cache {cache_name}: {e}") frs = { 'roadCameraState': FrameReader(get_url(TEST_ROUTE, SEGMENT, "fcamera.hevc"), pix_fmt='nv12', cache_size=END_FRAME - START_FRAME), 'driverCameraState': FrameReader(get_url(TEST_ROUTE, SEGMENT, "dcamera.hevc"), pix_fmt='nv12', cache_size=END_FRAME - START_FRAME), 'wideRoadCameraState': FrameReader(get_url(TEST_ROUTE, SEGMENT, "ecamera.hevc"), pix_fmt='nv12', cache_size=END_FRAME - START_FRAME), } for fr in frs.values(): for fidx in range(START_FRAME, END_FRAME): fr.get(fidx) fr.it = None print(f"Dumping frame cache {cache_name}") pickle.dump(frs, open(cache_name, "wb")) return frs if __name__ == "__main__": update = "--update" in sys.argv or (os.getenv("GIT_BRANCH", "") == 'master') replay_dir = os.path.dirname(os.path.abspath(__file__)) # load logs lr = list(LogReader(get_url(TEST_ROUTE, SEGMENT, "rlog.zst"))) frs = get_frames() log_msgs = [] # run replays log_msgs += model_replay(lr, frs) # get diff failed = False if not update: log_fn = get_log_fn(TEST_ROUTE) try: all_logs = list(LogReader(get_model_replay_url(log_fn))) cmp_log = [] model_start_index = next(i for i, m in enumerate(all_logs) if m.which() in ("modelV2", "drivingModelData", "cameraOdometry")) cmp_log += all_logs[model_start_index+START_FRAME*3:model_start_index + END_FRAME*3] dmon_start_index = next(i for i, m in enumerate(all_logs) if m.which() == "driverStateV2") cmp_log += all_logs[dmon_start_index+START_FRAME:dmon_start_index + END_FRAME] ignore = [ 'logMonoTime', 'drivingModelData.frameDropPerc', 'drivingModelData.modelExecutionTime', 'modelV2.frameDropPerc', 'modelV2.modelExecutionTime', 'driverStateV2.modelExecutionTime', 'driverStateV2.gpuExecutionTime' ] if PC: # TODO We ignore whole bunch so we can compare important stuff # like posenet with reasonable tolerance ignore += ['modelV2.acceleration.x', 'modelV2.position.x', 'modelV2.position.xStd', 'modelV2.position.y', 'modelV2.position.yStd', 'modelV2.position.z', 'modelV2.position.zStd', 'drivingModelData.path.xCoefficients',] for i in range(3): for field in ('x', 'y', 'v', 'a'): ignore.append(f'modelV2.leadsV3.{i}.{field}') ignore.append(f'modelV2.leadsV3.{i}.{field}Std') for i in range(4): for field in ('x', 'y', 'z', 't'): ignore.append(f'modelV2.laneLines.{i}.{field}') for i in range(2): for field in ('x', 'y', 'z', 't'): ignore.append(f'modelV2.roadEdges.{i}.{field}') tolerance = .3 if PC else None results: Any = {TEST_ROUTE: {}} log_paths: Any = {TEST_ROUTE: {"models": {'ref': log_fn, 'new': log_fn}}} results[TEST_ROUTE]["models"] = compare_logs(cmp_log, log_msgs, tolerance=tolerance, ignore_fields=ignore) diff_short, diff_long, failed = format_diff(results, log_paths, 'master') if "CI" in os.environ: failed = False print(diff_long) print('-------------\n'*5) print(diff_short) with open("model_diff.txt", "w") as f: f.write(diff_long) except Exception as e: print(str(e)) failed = True if update and not PC: log_fn = get_log_fn(TEST_ROUTE) save_log(log_fn, log_msgs) sys.exit(int(failed))