diff --git a/opendbc_repo/opendbc/car/hyundai/carcontroller.py b/opendbc_repo/opendbc/car/hyundai/carcontroller.py index 1fa8bd0a1..99f6b3807 100644 --- a/opendbc_repo/opendbc/car/hyundai/carcontroller.py +++ b/opendbc_repo/opendbc/car/hyundai/carcontroller.py @@ -158,6 +158,15 @@ def should_track_stop_accel_directly(stopping: bool, v_ego: float, return bool(stopping and v_ego > EV6_GT_LINE_STOP_BRAKE_CAP_MAX_SPEED and accel_cmd < actual_accel) +def should_track_stop_accel_directly_for_car(car_fingerprint, stopping: bool, v_ego: float, + accel_cmd: float, actual_accel: float) -> bool: + # EV9 uses the low-speed stop cap to avoid a harsh lead re-brake handoff. + # Keep direct tracking for the other CCNC angle-steering platform. + return car_fingerprint != CAR.KIA_EV9 and should_track_stop_accel_directly( + stopping, v_ego, accel_cmd, actual_accel, + ) + + def should_use_ev6_gt_line_stop_direct_tracking(ev6_gt_line: bool, stopping: bool, v_ego: float, accel_cmd: float, actual_accel: float) -> bool: return bool(ev6_gt_line and stopping and v_ego > EV6_GT_LINE_STOP_BRAKE_CAP_MAX_SPEED and accel_cmd < actual_accel) @@ -399,6 +408,11 @@ def process_hud_alert(enabled, fingerprint, hud_control): return sys_warning, sys_state, left_lane_warning, right_lane_warning +def preserve_stock_canfd_lfa_status(car_fingerprint) -> bool: + # The 2022-24 Carnival expects a clean replacement status payload after its radar ECU is disabled. + return car_fingerprint != CAR.KIA_CARNIVAL_4TH_GEN + + class CarController(CarControllerBase): def __init__(self, dbc_names, CP): super().__init__(dbc_names, CP) @@ -630,8 +644,10 @@ class CarController(CarControllerBase): if should_use_ev6_gt_line_stop_direct_tracking(is_ev6_gt_line, self._ioniq_6_long_tuning.stopping, CS.out.vEgo, accel_cmd, self._ioniq_6_long_tuning.actual_accel): use_egmp_smoothed_accel = False - if is_ccnc_angle_long and should_track_stop_accel_directly(self._ioniq_6_long_tuning.stopping, CS.out.vEgo, - accel_cmd, self._ioniq_6_long_tuning.actual_accel): + if is_ccnc_angle_long and should_track_stop_accel_directly_for_car( + self.CP.carFingerprint, self._ioniq_6_long_tuning.stopping, CS.out.vEgo, + accel_cmd, self._ioniq_6_long_tuning.actual_accel, + ): use_egmp_smoothed_accel = False if use_egmp_dynamic_long_tuning: if use_egmp_smoothed_accel: @@ -803,10 +819,11 @@ class CarController(CarControllerBase): forward_stock_lkas = angle_lkas_alt and ( angle_lkas_alt_standstill_handoff or not (drive_gear and (CC.latActive or CC.enabled)) ) + preserve_stock_lfa_status = preserve_stock_canfd_lfa_status(self.CP.carFingerprint) if not forward_stock_lkas and not ccnc_angle_long: can_sends.extend(hyundaicanfd.create_steering_messages(self.packer, self.CP, self.CAN, CC.enabled, steering_msg_active, apply_torque, apply_angle, - CS.stock_lfa_msg, + CS.stock_lfa_msg if preserve_stock_lfa_status else None, CS.stock_lkas_msg if preserve_stock_lkas else None, lka_icon=lka_icon, send_lfa_status=self.ecu_disable_failed and @@ -845,7 +862,8 @@ class CarController(CarControllerBase): CC.leftBlinker, CC.rightBlinker, CS.msg_161, CS.msg_162, CS.msg_1b5, CS.is_metric, CS.out, CS.out.cruiseState.available, lfa_icon)) else: - can_sends.append(hyundaicanfd.create_lfahda_cluster(self.packer, self.CAN, CC.enabled, CS.stock_lfahda_cluster_msg, + cluster_base_values = CS.stock_lfahda_cluster_msg if preserve_stock_lfa_status else None + can_sends.append(hyundaicanfd.create_lfahda_cluster(self.packer, self.CAN, CC.enabled, cluster_base_values, lfa_icon=lfa_icon)) # blinkers diff --git a/opendbc_repo/opendbc/car/hyundai/tests/test_hyundai.py b/opendbc_repo/opendbc/car/hyundai/tests/test_hyundai.py index a3a9862b3..ab89bab05 100644 --- a/opendbc_repo/opendbc/car/hyundai/tests/test_hyundai.py +++ b/opendbc_repo/opendbc/car/hyundai/tests/test_hyundai.py @@ -16,7 +16,9 @@ from opendbc.car.hyundai.carcontroller import CarController, Ioniq6LongitudinalT get_canfd_scc_decel_step, \ should_reset_ev6_gt_line_longitudinal_tuning, reset_ev6_gt_line_longitudinal_tuning, \ direct_angle_request_allowed, get_angle_smoothing_alpha, \ - should_use_ev6_gt_line_stop_direct_tracking + should_use_ev6_gt_line_stop_direct_tracking, \ + should_track_stop_accel_directly_for_car, \ + preserve_stock_canfd_lfa_status from opendbc.car.hyundai.carstate import CarState, decode_canfd_camera_lead, decode_ioniq_6_blindspot_radar_state, \ get_canfd_cruise_available from opendbc.car.hyundai.interface import CarInterface, KIA_EV9_ACCEL_MAX @@ -122,6 +124,29 @@ def get_test_toggles() -> SimpleNamespace: class TestHyundaiFingerprint: + def test_carnival_2024_uses_clean_canfd_lfa_status(self): + assert not preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_4TH_GEN) + assert preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_2025) + assert preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_HEV_4TH_GEN) + assert preserve_stock_canfd_lfa_status(CAR.HYUNDAI_IONIQ_6) + + CP = CarParams.new_message() + CP.carFingerprint = CAR.KIA_CARNIVAL_4TH_GEN + CP.flags = int(HyundaiFlags.CANFD | HyundaiFlags.RADAR_SCC) + CP.openpilotLongitudinalControl = True + packer = CANPacker(DBC[CP.carFingerprint][Bus.pt]) + can_bus = CanBus(CP) + + stock_lfa = {"HAS_LANE_SAFETY": 1, "NEW_SIGNAL_4": 8, "DAMP_FACTOR": 100} + lfa_base = stock_lfa if preserve_stock_canfd_lfa_status(CP.carFingerprint) else None + lfa_msg = hyundaicanfd.create_steering_messages(packer, CP, can_bus, False, False, 0, 0.0, lfa_base)[0] + assert lfa_msg[1] == bytes.fromhex("05100002400008000000000000640000") + + stock_cluster = {"NEW_SIGNAL_5": 1} + cluster_base = stock_cluster if preserve_stock_canfd_lfa_status(CP.carFingerprint) else None + cluster_msg = hyundaicanfd.create_lfahda_cluster(packer, can_bus, False, cluster_base) + assert cluster_msg[1] == bytes.fromhex("8e040000000000000000000000000000") + def test_canfd_torque_bsm_parser_registers_rear_blindspots(self): CP = CarParams.new_message() CP.carFingerprint = CAR.GENESIS_GV70_ELECTRIFIED_1ST_GEN @@ -1417,6 +1442,14 @@ class TestHyundaiFingerprint: assert get_canfd_scc_decel_step(ev9_cp) == pytest.approx(0.20) assert get_canfd_scc_decel_step(ioniq_6_cp) == pytest.approx(0.36) + def test_ev9_keeps_low_speed_stop_brake_cap(self): + assert not should_track_stop_accel_directly_for_car( + CAR.KIA_EV9, stopping=True, v_ego=1.8, accel_cmd=-3.5, actual_accel=-0.4, + ) + assert should_track_stop_accel_directly_for_car( + CAR.HYUNDAI_IONIQ_5_PE, stopping=True, v_ego=1.8, accel_cmd=-3.5, actual_accel=-0.4, + ) + def test_ev9_longitudinal_decel_jerk_is_bounded(self): state = Ioniq6LongitudinalTuningState(actual_accel=-1.0, accel_last=-1.0) state = update_ioniq_6_longitudinal_tuning( diff --git a/opendbc_repo/opendbc/car/subaru/carcontroller.py b/opendbc_repo/opendbc/car/subaru/carcontroller.py index f0c35416b..c9db33497 100644 --- a/opendbc_repo/opendbc/car/subaru/carcontroller.py +++ b/opendbc_repo/opendbc/car/subaru/carcontroller.py @@ -160,7 +160,7 @@ class CarController(CarControllerBase): self.CP.openpilotLongitudinalControl, CC.longActive, hud_control.leadVisible, self.status_bus)) - can_sends.append(subarucan.create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.enabled, hud_control.visualAlert, + can_sends.append(subarucan.create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.latActive, hud_control.visualAlert, hud_control.leftLaneVisible, hud_control.rightLaneVisible, hud_control.leftLaneDepart, hud_control.rightLaneDepart, self.status_bus)) diff --git a/opendbc_repo/opendbc/car/subaru/subarucan.py b/opendbc_repo/opendbc/car/subaru/subarucan.py index 93a354ab2..f52a1a190 100644 --- a/opendbc_repo/opendbc/car/subaru/subarucan.py +++ b/opendbc_repo/opendbc/car/subaru/subarucan.py @@ -124,6 +124,7 @@ def create_es_lkas_state(packer, frame, es_lkas_state_msg, enabled, visual_alert values["LKAS_ACTIVE"] = 1 # Show LKAS lane lines values["LKAS_Dash_State"] = 2 # Green enabled indicator else: + values["LKAS_ACTIVE"] = 0 values["LKAS_Dash_State"] = 0 # LKAS Not enabled values["LKAS_Left_Line_Visible"] = int(left_line) diff --git a/opendbc_repo/opendbc/car/subaru/tests/test_subaru.py b/opendbc_repo/opendbc/car/subaru/tests/test_subaru.py index bdbaf01dd..7f13f583e 100644 --- a/opendbc_repo/opendbc/car/subaru/tests/test_subaru.py +++ b/opendbc_repo/opendbc/car/subaru/tests/test_subaru.py @@ -1,14 +1,18 @@ +import inspect +from collections import defaultdict from types import SimpleNamespace import pytest +from opendbc.can import CANPacker, CANParser from opendbc.car import Bus +from opendbc.car.subaru import subarucan from opendbc.car.subaru.carcontroller import CarController from opendbc.car.subaru.carstate import CarState from opendbc.car.subaru.fingerprints import FW_VERSIONS from opendbc.car.fw_versions import match_fw_to_car from opendbc.car.subaru.interface import CarInterface -from opendbc.car.subaru.values import CAR, CanBus, SubaruFlags, SubaruSafetyFlags +from opendbc.car.subaru.values import CAR, DBC, CanBus, SubaruFlags, SubaruSafetyFlags from opendbc.car.structs import CarParams @@ -90,7 +94,6 @@ class TestSubaruFingerprint: assert matches == {CAR.SUBARU_OUTBACK_2023} def test_legacy_2025_firmware(self): - legacy_fw = FW_VERSIONS[CAR.SUBARU_LEGACY_2025] car_fw = [ CarParams.CarFw(ecu=CarParams.Ecu.abs, fwVersion=b'\xa1 $\x11\x00', address=0x7b0, brand="subaru"), CarParams.CarFw(ecu=CarParams.Ecu.eps, fwVersion=b'[\xc0\xd1\x10\x00', address=0x746, brand="subaru"), @@ -172,23 +175,23 @@ def test_outback_2023_uses_d_platform_bus_layout(): assert controller.status_bus == CanBus.main -def test_legacy_2025_uses_d_platform_bus_layout(): +def test_legacy_2025_uses_gen2_angle_bus_layout(): CP = CarInterface.get_non_essential_params(CAR.SUBARU_LEGACY_2025) parsers = CarState.get_can_parsers(CP) controller = CarController({}, CP) - assert CP.flags & SubaruFlags.D_PLATFORM - assert CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM - assert CP.flags & SubaruFlags.D_PLATFORM_CAMERA - assert CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM_CAMERA - assert CanBus.main_for_cp(CP) == CanBus.alt - assert CanBus.angle_for_cp(CP) == CanBus.camera - assert parsers[Bus.pt].bus == CanBus.alt + assert not (CP.flags & SubaruFlags.D_PLATFORM) + assert not (CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM) + assert not (CP.flags & SubaruFlags.D_PLATFORM_CAMERA) + assert not (CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM_CAMERA) + assert CanBus.main_for_cp(CP) == CanBus.main + assert CanBus.angle_for_cp(CP) == CanBus.main + assert parsers[Bus.pt].bus == CanBus.main assert parsers[Bus.cam].bus == CanBus.camera assert parsers[Bus.alt].bus == CanBus.alt - assert parsers[Bus.main].bus == CanBus.main - assert controller.angle_bus == CanBus.camera - assert controller.status_bus == CanBus.camera + assert Bus.main not in parsers + assert controller.angle_bus == CanBus.main + assert controller.status_bus == CanBus.main def test_ascent_2023_uses_d_platform_bus_layout(): @@ -232,3 +235,26 @@ def test_angle_controller_tracks_driver_override(): assert controller.driver_override assert controller.apply_steer_last == CS.out.steeringAngleDeg assert msg[0] == 0x124 + + +def test_lkas_hud_state_uses_lateral_active(): + update_source = inspect.getsource(CarController.update) + + assert "create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.latActive" in update_source + assert "create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.enabled" not in update_source + + +@pytest.mark.parametrize(("enabled", "expected"), ((False, 0), (True, 1))) +def test_lkas_hud_active_bit_follows_lateral_state(enabled, expected): + dbc = DBC[CAR.SUBARU_LEGACY_2025][Bus.pt] + packer = CANPacker(dbc) + parser = CANParser(dbc, [("ES_LKAS_State", 0)], CanBus.main) + stock_lkas_state = defaultdict(int, {"LKAS_ACTIVE": 1}) + + msg = subarucan.create_es_lkas_state( + packer, 0, stock_lkas_state, enabled, 0, False, False, False, False, CanBus.main, + ) + parser.update([(1, [msg])]) + + assert parser.can_valid + assert parser.vl["ES_LKAS_State"]["LKAS_ACTIVE"] == expected diff --git a/opendbc_repo/opendbc/car/subaru/values.py b/opendbc_repo/opendbc/car/subaru/values.py index 1aac1b7fc..15116d0d6 100644 --- a/opendbc_repo/opendbc/car/subaru/values.py +++ b/opendbc_repo/opendbc/car/subaru/values.py @@ -242,7 +242,7 @@ class CAR(Platforms): SUBARU_LEGACY_2025 = SubaruGen2PlatformConfig( [SubaruCarDocs("Subaru Legacy 2025", "All", car_parts=CarParts.common([CarHarness.subaru_d]))], SUBARU_OUTBACK.specs, - flags=SubaruFlags.LKAS_ANGLE | SubaruFlags.D_PLATFORM | SubaruFlags.D_PLATFORM_CAMERA, + flags=SubaruFlags.LKAS_ANGLE, ) SUBARU_ASCENT_2023 = SubaruGen2PlatformConfig( [SubaruCarDocs("Subaru Ascent 2023-25", "All", car_parts=CarParts.common([CarHarness.subaru_d]))], diff --git a/opendbc_repo/opendbc/safety/modes/subaru.h b/opendbc_repo/opendbc/safety/modes/subaru.h index df92fb176..dd1f4c26d 100644 --- a/opendbc_repo/opendbc/safety/modes/subaru.h +++ b/opendbc_repo/opendbc/safety/modes/subaru.h @@ -90,6 +90,7 @@ {.msg = {{MSG_SUBARU_Wheel_Speeds, alt_bus, 8, 50U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \ {.msg = {{MSG_SUBARU_Brake_Status, alt_bus, 8, 50U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \ {.msg = {{MSG_SUBARU_ES_Status, status_bus, 8, 20U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \ + {.msg = {{MSG_SUBARU_ES_DashStatus, SUBARU_CAM_BUS, 8, 10U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \ #define SUBARU_D_PLATFORM_ANGLE_RX_CHECKS() \ {.msg = {{MSG_SUBARU_Throttle, SUBARU_ALT_BUS, 8, 100U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \ @@ -148,6 +149,10 @@ static void subaru_rx_hook(const CANPacket_t *msg) { pcm_cruise_check(cruise_engaged); } + if (subaru_lkas_angle && (msg->addr == MSG_SUBARU_ES_DashStatus) && (msg->bus == SUBARU_CAM_BUS)) { + acc_main_on = GET_BIT(msg, 49U); + } + if (!subaru_lkas_angle && (msg->addr == MSG_SUBARU_CruiseControl) && (msg->bus == alt_main_bus)) { bool cruise_engaged = (msg->data[5] >> 1) & 1U; pcm_cruise_check(cruise_engaged); diff --git a/opendbc_repo/opendbc/safety/tests/test_subaru.py b/opendbc_repo/opendbc/safety/tests/test_subaru.py index bb5b3fbca..cf9b6695a 100755 --- a/opendbc_repo/opendbc/safety/tests/test_subaru.py +++ b/opendbc_repo/opendbc/safety/tests/test_subaru.py @@ -9,6 +9,7 @@ from opendbc.car.subaru.carcontroller import get_safety_CP from opendbc.car.subaru.values import CarControllerParams, SubaruSafetyFlags from opendbc.car.structs import CarParams from opendbc.car.vehicle_model import VehicleModel +from opendbc.safety import ALTERNATIVE_EXPERIENCE from opendbc.safety.tests.libsafety import libsafety_py import opendbc.safety.tests.common as common from opendbc.safety.tests.common import CANPackerSafety, away_round, round_speed @@ -230,6 +231,21 @@ class TestSubaruAngleSafetyBase(TestSubaruSafetyBase, common.AngleSteeringSafety def _toggle_aol(self, toggle_on): return None + def _acc_main_msg(self, main_on): + values = {"Cruise_On": int(main_on)} + return self.packer.make_can_msg_panda("ES_DashStatus", SUBARU_CAM_BUS, values) + + def test_acc_main_tracks_dash_status(self): + self.safety.set_alternative_experience(ALTERNATIVE_EXPERIENCE.ALWAYS_ON_LATERAL) + + self._rx(self._acc_main_msg(False)) + self.assertFalse(self.safety.get_acc_main_on()) + self.assertFalse(self.safety.get_aol_allowed()) + + self._rx(self._acc_main_msg(True)) + self.assertTrue(self.safety.get_acc_main_on()) + self.assertTrue(self.safety.get_aol_allowed()) + def test_angle_cmd_when_enabled(self): pass diff --git a/selfdrive/controls/lib/latcontrol_torque.py b/selfdrive/controls/lib/latcontrol_torque.py index b0c86b4a2..9799ca2eb 100644 --- a/selfdrive/controls/lib/latcontrol_torque.py +++ b/selfdrive/controls/lib/latcontrol_torque.py @@ -558,6 +558,9 @@ class LatControlTorque(LatControl): elif genesis_g70_active: output_torque *= genesis_g70_center_output_taper output_torque *= get_genesis_g70_curve_unwind_output_scale(setpoint, desired_lateral_jerk, CS.vEgo) + output_torque *= get_genesis_g70_high_speed_error_scale( + setpoint, measurement, desired_lateral_jerk, CS.vEgo, + ) low_speed_output_limit = get_genesis_g70_low_speed_output_limit(setpoint, CS.vEgo) output_torque = float(np.clip(output_torque, -low_speed_output_limit, low_speed_output_limit)) elif sonata_hybrid_active: diff --git a/selfdrive/controls/lib/latcontrol_vehicle_tunes.py b/selfdrive/controls/lib/latcontrol_vehicle_tunes.py index b32916c40..308ac8e68 100644 --- a/selfdrive/controls/lib/latcontrol_vehicle_tunes.py +++ b/selfdrive/controls/lib/latcontrol_vehicle_tunes.py @@ -252,6 +252,13 @@ GENESIS_G70_CURVE_UNWIND_LAT = 0.25 GENESIS_G70_CURVE_UNWIND_LAT_WIDTH = 0.12 GENESIS_G70_CURVE_UNWIND_JERK = 0.08 GENESIS_G70_CURVE_UNWIND_JERK_WIDTH = 0.08 +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_MAX = 0.15 +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED = 50.0 * CV.MPH_TO_MS +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED_WIDTH = 8.0 * CV.MPH_TO_MS +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR = 0.18 +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR_WIDTH = 0.15 +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK = 0.15 +GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK_WIDTH = 0.10 BOLT_2017_LATERAL_TESTING_GROUND_ID = testing_ground.id_3 BOLT_2017_STEER_RATIO_TEST_SCALE = 1.045 @@ -321,12 +328,14 @@ BOLT_2022_2023_CENTER_TAPER_LAT = 0.18 BOLT_2022_2023_CENTER_TAPER_LAT_WIDTH = 0.03 BOLT_2022_2023_CENTER_TAPER_SPEED = 25.0 BOLT_2022_2023_CENTER_TAPER_SPEED_WIDTH = 2.5 -BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX = 0.07 +BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX = 0.12 BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT = 0.14 BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT_WIDTH = 0.04 BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED = 4.0 BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_WIDTH = 1.5 -BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX = 14.0 +BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR = 2.0 +BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR_WIDTH = 0.7 +BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX = 16.5 BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX_WIDTH = 2.0 BOLT_2022_2023_LOW_SPEED_CENTER_OUTPUT_LIMIT = 0.38 BOLT_2022_2023_LOW_SPEED_CENTER_OUTPUT_LAT = 0.17 @@ -1903,10 +1912,14 @@ def get_bolt_2022_2023_center_output_scale(desired_lateral_accel: float, v_ego: BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX_WIDTH) low_speed_center_weight = _bolt_2022_2023_sigmoid((BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT - abs(desired_lateral_accel)) / BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT_WIDTH) + low_speed_floor = _bolt_2022_2023_sigmoid( + (v_ego - BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR) / + BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR_WIDTH + ) highway_reduction = (_flm_vehicle_knob("gm_bolt_2022_2023.center_taper_max", BOLT_2022_2023_CENTER_TAPER_MAX) * highway_speed_weight * highway_center_weight) low_speed_reduction = (BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX * low_speed_onset * low_speed_cutoff * - low_speed_center_weight) + low_speed_center_weight * low_speed_floor) return 1.0 - min(highway_reduction + low_speed_reduction, 0.95) @@ -2685,6 +2698,23 @@ def get_genesis_g70_curve_unwind_output_scale(desired_lateral_accel: float, desi return 1.0 + GENESIS_G70_CURVE_UNWIND_OUTPUT_BOOST * speed_weight * lateral_weight * jerk_weight +def get_genesis_g70_high_speed_error_scale(setpoint: float, measured_lateral_accel: float, + desired_lateral_jerk: float, v_ego: float) -> float: + tracking_error = abs(measured_lateral_accel - setpoint) + if tracking_error <= 0.0: + return 1.0 + speed_weight = _sigmoid((v_ego - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED) / + GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED_WIDTH) + error_weight = _sigmoid((tracking_error - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR) / + GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR_WIDTH) + jerk_weight = _sigmoid((abs(desired_lateral_jerk) - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK) / + GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK_WIDTH) + phase_weight = 1.0 if setpoint * desired_lateral_jerk < 0.0 else 0.45 + reduction = (GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_MAX * speed_weight * error_weight * + (0.35 + (0.65 * jerk_weight)) * phase_weight) + return 1.0 - reduction + + def _ioniq_5_sigmoid(x: float) -> float: return _sigmoid(x) diff --git a/selfdrive/controls/lib/longitudinal_vehicle_tunes.py b/selfdrive/controls/lib/longitudinal_vehicle_tunes.py index 198f62985..b428adae6 100644 --- a/selfdrive/controls/lib/longitudinal_vehicle_tunes.py +++ b/selfdrive/controls/lib/longitudinal_vehicle_tunes.py @@ -4,6 +4,7 @@ import numpy as np HONDA_HRV_3G_FAR_FOLLOW_BRAKE_SLEW_RATE = 3.0 HONDA_HRV_3G_FAR_FOLLOW_RELEASE_SLEW_RATE = 2.0 HONDA_HRV_3G_UNTRACKED_SLOW_LEAD_DECEL_SCALE = 1.35 +HYUNDAI_ELANTRA_LEAD_FOLLOW_JERK_SCALE = 1.25 GM_SILVERADO_EARLY_FOLLOW_MIN_EGO_SPEED = 18.0 GM_SILVERADO_EARLY_FOLLOW_MAX_DISTANCE = 130.0 GM_SILVERADO_EARLY_FOLLOW_MIN_MODEL_PROB = 0.85 @@ -42,6 +43,13 @@ def get_untracked_slow_lead_decel_scale(CP): return 1.0 +def get_lead_follow_jerk_scale(CP): + """Spread the lead-source transition for cars with a sharp vision-lead handoff.""" + if getattr(CP, "brand", "") == "hyundai" and str(getattr(CP, "carFingerprint", "")) == "HYUNDAI_ELANTRA_2021": + return HYUNDAI_ELANTRA_LEAD_FOLLOW_JERK_SCALE + return 1.0 + + def is_gm_silverado_early_follow_lead(CP, lead, v_ego): """Admit a credible centered vision lead before it becomes a close lead.""" if ( diff --git a/selfdrive/controls/tests/test_latcontrol.py b/selfdrive/controls/tests/test_latcontrol.py index bbd3a1924..d7a210e84 100644 --- a/selfdrive/controls/tests/test_latcontrol.py +++ b/selfdrive/controls/tests/test_latcontrol.py @@ -74,6 +74,7 @@ from openpilot.selfdrive.controls.lib.latcontrol_torque import ( get_genesis_g70_curve_unwind_output_scale, get_genesis_g70_friction_jerk_deadzone, get_genesis_g70_friction_threshold, + get_genesis_g70_high_speed_error_scale, get_genesis_g70_low_speed_angle_damping, get_genesis_g70_low_speed_output_limit, get_genesis_gv70_friction_threshold, @@ -286,9 +287,9 @@ class TestLatControl: highway_turn = get_bolt_2022_2023_center_output_scale(0.40, 31.0) creep_center = get_bolt_2022_2023_center_output_scale(0.04, 1.0) - assert 0.92 < low_speed_center < 0.95 + assert 0.88 < low_speed_center < 0.92 assert low_speed_turn > 0.99 - assert middle_speed_center > 0.98 + assert 0.96 < middle_speed_center < 0.98 assert 0.88 < highway_center < 0.91 assert highway_turn > 0.99 assert creep_center > 0.99 @@ -812,6 +813,10 @@ class TestLatControl: assert get_genesis_g70_curve_unwind_output_scale(0.7, -0.5, 25.0) > 1.0 assert get_genesis_g70_curve_unwind_output_scale(0.7, 0.5, 25.0) == 1.0 + assert get_genesis_g70_high_speed_error_scale(0.2, 0.2, 0.8, 20.0) == 1.0 + assert get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 20.0) < 1.0 + assert get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 10.0) > get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 20.0) + def test_sonata_hybrid_center_output_taper_is_mid_speed_and_center_gated(self): low_speed = get_sonata_hybrid_center_output_scale(0.0, 8.0) center = get_sonata_hybrid_center_output_scale(0.0, 13.4) diff --git a/selfdrive/controls/tests/test_starpilot_planner.py b/selfdrive/controls/tests/test_starpilot_planner.py index 4ce8d2bea..b70e427f1 100644 --- a/selfdrive/controls/tests/test_starpilot_planner.py +++ b/selfdrive/controls/tests/test_starpilot_planner.py @@ -4,6 +4,7 @@ from types import SimpleNamespace from openpilot.common.realtime import DT_MDL from openpilot.starpilot.controls.starpilot_planner import StarPilotPlanner, get_force_stop_jerk_scale +from openpilot.selfdrive.controls.lib.longitudinal_vehicle_tunes import get_lead_follow_jerk_scale import openpilot.starpilot.controls.starpilot_planner as starpilot_planner_module @@ -38,6 +39,11 @@ def test_force_stop_jerk_scale_is_platform_specific(): assert get_force_stop_jerk_scale(SimpleNamespace(carFingerprint="OTHER_CAR")) == 0.32 +def test_lead_follow_jerk_scale_is_platform_specific(): + assert get_lead_follow_jerk_scale(SimpleNamespace(brand="hyundai", carFingerprint="HYUNDAI_ELANTRA_2021")) == 1.25 + assert get_lead_follow_jerk_scale(SimpleNamespace(brand="other", carFingerprint="OTHER_CAR")) == 1.0 + + def make_sm(planner, *, frame: int, v_ego: float, left_blinker: bool, right_blinker: bool = False, standstill: bool = False): return FakeSM(frame, { "radarState": SimpleNamespace( diff --git a/selfdrive/modeld/get_model_metadata.py b/selfdrive/modeld/get_model_metadata.py index 40a8c0c45..e4c173957 100755 --- a/selfdrive/modeld/get_model_metadata.py +++ b/selfdrive/modeld/get_model_metadata.py @@ -1,42 +1,55 @@ #!/usr/bin/env python3 +import sys +import pathlib import codecs import pickle -import pathlib -import sys from typing import Any -import onnx +from tinygrad.nn.onnx import OnnxPBParser -def get_name_and_shape(value_info:onnx.ValueInfoProto) -> tuple[str, tuple[int,...]]: - shape = tuple([int(dim.dim_value) for dim in value_info.type.tensor_type.shape.dim]) - name = value_info.name + +class MetadataOnnxPBParser(OnnxPBParser): + def _parse_ModelProto(self) -> dict: + obj: dict[str, Any] = {"graph": {"input": [], "output": []}, "metadata_props": []} + for fid, wire_type in self._parse_message(self.reader.len): + match fid: + case 7: + obj["graph"] = self._parse_GraphProto() + case 14: + obj["metadata_props"].append(self._parse_StringStringEntryProto()) + case _: + self.reader.skip_field(wire_type) + return obj + + +def get_name_and_shape(value_info: dict[str, Any]) -> tuple[str, tuple[int, ...]]: + shape = tuple(int(dim) if isinstance(dim, int) else 0 for dim in value_info["parsed_type"].shape) + name = value_info["name"] return name, shape -def get_metadata_value_by_name(model:onnx.ModelProto, name:str) -> str | Any: - for prop in model.metadata_props: - if prop.key == name: - return prop.value + +def get_metadata_value_by_name(model: dict[str, Any], name: str) -> str | Any: + for prop in model["metadata_props"]: + if prop["key"] == name: + return prop["value"] return None -def make_metadata_dict(model_path: str | pathlib.Path) -> dict[str, Any]: - model = onnx.load(str(model_path)) +def make_metadata_dict(model_path): + model = MetadataOnnxPBParser(model_path).parse() output_slices = get_metadata_value_by_name(model, 'output_slices') assert output_slices is not None, 'output_slices not found in metadata' - return { 'model_checkpoint': get_metadata_value_by_name(model, 'model_checkpoint'), 'output_slices': pickle.loads(codecs.decode(output_slices.encode(), "base64")), - 'input_shapes': dict([get_name_and_shape(x) for x in model.graph.input]), - 'output_shapes': dict([get_name_and_shape(x) for x in model.graph.output]), + 'input_shapes': dict(get_name_and_shape(x) for x in model["graph"]["input"]), + 'output_shapes': dict(get_name_and_shape(x) for x in model["graph"]["output"]), } + if __name__ == "__main__": model_path = pathlib.Path(sys.argv[1]) - metadata = make_metadata_dict(model_path) - metadata_path = model_path.parent / (model_path.stem + '_metadata.pkl') with open(metadata_path, 'wb') as f: - pickle.dump(metadata, f) - + pickle.dump(make_metadata_dict(model_path), f) print(f'saved metadata to {metadata_path}') diff --git a/selfdrive/modeld/tests/test_get_model_metadata.py b/selfdrive/modeld/tests/test_get_model_metadata.py new file mode 100644 index 000000000..49907da56 --- /dev/null +++ b/selfdrive/modeld/tests/test_get_model_metadata.py @@ -0,0 +1,35 @@ +import base64 +import pickle + +import onnx + +from openpilot.selfdrive.modeld.get_model_metadata import make_metadata_dict + + +def test_make_metadata_dict_uses_disk_backed_parser(tmp_path): + model_input = onnx.helper.make_tensor_value_info("input", onnx.TensorProto.FLOAT, (1, 3)) + model_output = onnx.helper.make_tensor_value_info("output", onnx.TensorProto.FLOAT, (1, 4)) + graph = onnx.helper.make_graph( + [onnx.helper.make_node("Identity", ["input"], ["output"])], + "metadata-test", + [model_input], + [model_output], + ) + model = onnx.helper.make_model(graph) + output_slices = {"path": slice(0, 4)} + model.metadata_props.append(onnx.StringStringEntryProto( + key="output_slices", + value=base64.b64encode(pickle.dumps(output_slices)).decode(), + )) + model.metadata_props.append(onnx.StringStringEntryProto(key="model_checkpoint", value="test-checkpoint")) + model_path = tmp_path / "model.onnx" + onnx.save(model, model_path) + + metadata = make_metadata_dict(model_path) + + assert metadata == { + "model_checkpoint": "test-checkpoint", + "output_slices": output_slices, + "input_shapes": {"input": (1, 3)}, + "output_shapes": {"output": (1, 4)}, + } diff --git a/selfdrive/ui/mici/layouts/settings/developer.py b/selfdrive/ui/mici/layouts/settings/developer.py index ee18c46b9..1a396a7ed 100644 --- a/selfdrive/ui/mici/layouts/settings/developer.py +++ b/selfdrive/ui/mici/layouts/settings/developer.py @@ -119,6 +119,13 @@ class DeveloperLayoutMici(NavScroller): super()._update_state() self._ssh_fetcher.update() + def show_event(self): + super().show_event() + # CarParamsPersistent can change while this panel is not visible (for + # example after applying a manual fingerprint). Refresh capability-gated + # controls when the panel is opened so stale controls are not shown. + self._update_toggles() + def _update_toggles(self): ui_state.update_params() diff --git a/starpilot/assets/tests/test_model_pipeline.py b/starpilot/assets/tests/test_model_pipeline.py index 34e6ada51..a9ce3fc97 100644 --- a/starpilot/assets/tests/test_model_pipeline.py +++ b/starpilot/assets/tests/test_model_pipeline.py @@ -66,7 +66,7 @@ def test_external_gpu_requirement_is_cached_from_manifest(tmp_path, monkeypatch) def test_external_gpu_compilation_is_opt_in(tmp_path, monkeypatch): invocations = [] - monkeypatch.setattr(model_compiler, "build_compile_env", lambda: { + monkeypatch.setattr(model_compiler, "build_compile_env", lambda **_: { "DEV": "QCOM", "IMAGE": "2", "NOLOCALS": "1", "OPENPILOT_HACKS": "1", }) monkeypatch.setattr(model_compiler.subprocess, "run", lambda command, **kwargs: invocations.append((command, kwargs))) @@ -88,7 +88,7 @@ def test_external_gpu_compilation_is_opt_in(tmp_path, monkeypatch): def test_compile_clears_only_selected_model_outputs(tmp_path, monkeypatch): - monkeypatch.setattr(model_compiler, "build_compile_env", lambda: {}) + monkeypatch.setattr(model_compiler, "build_compile_env", lambda **_: {}) monkeypatch.setattr(model_compiler.subprocess, "run", lambda *args, **kwargs: None) (tmp_path / "normal_driving_tinygrad.pkl").write_bytes(b"old") (tmp_path / "normal_driving_tinygrad.pkl.p00").write_bytes(b"old") diff --git a/starpilot/controls/starpilot_planner.py b/starpilot/controls/starpilot_planner.py index f54e0494c..78edbd1a0 100644 --- a/starpilot/controls/starpilot_planner.py +++ b/starpilot/controls/starpilot_planner.py @@ -18,6 +18,7 @@ from openpilot.selfdrive.controls.lib.lead_behavior import ( should_track_lead, ) from openpilot.selfdrive.controls.lib.longitudinal_mpc_lib.long_mpc import A_CHANGE_COST, DANGER_ZONE_COST, J_EGO_COST, STOP_DISTANCE +from openpilot.selfdrive.controls.lib.longitudinal_vehicle_tunes import get_lead_follow_jerk_scale from openpilot.starpilot.common.starpilot_utilities import calculate_lane_width, calculate_road_curvature from openpilot.starpilot.common.starpilot_variables import CRUISING_SPEED, MINIMUM_LATERAL_ACCELERATION, PLANNER_TIME, THRESHOLD @@ -300,12 +301,19 @@ class StarPilotPlanner: # While committed to a Force Stop, cut the MPC's accel-change penalty so terminal # braking can ramp faster. 0.32 lands near 40, what long_mpc uses in blended mode. + try: + car_params = sm["carParams"] + except (KeyError, IndexError, TypeError, AttributeError): + car_params = None + if self.starpilot_vcruise.forcing_stop: - try: - car_params = sm["carParams"] - except (KeyError, IndexError, TypeError, AttributeError): - car_params = None jerk_scale = get_force_stop_jerk_scale(car_params) + elif self.tracking_lead: + # Elantra vision leads can hand off from cruise to lead0 while closing + # quickly. A slightly higher accel-change cost makes that handoff begin + # earlier instead of arriving as a sharp brake request, without changing + # the safety stop distance or the force-stop path. + jerk_scale = get_lead_follow_jerk_scale(car_params) else: jerk_scale = 1.0 starpilotPlan.accelerationJerk = float(A_CHANGE_COST * self.starpilot_following.acceleration_jerk * jerk_scale) diff --git a/system/ui/README.md b/system/ui/README.md index 696cd4f63..3c42622ad 100644 --- a/system/ui/README.md +++ b/system/ui/README.md @@ -8,9 +8,10 @@ Quick start: * set `STRICT_MODE=1` to kill the app if it drops too much below 60fps * set `SCALE=1.5` to scale the entire UI by 1.5x * set `BURN_IN=1` to get a burn-in heatmap version of the UI -* burn-in prevention shifts the UI by 2 pixels every 3 minutes on device; set `BURN_IN_PREVENTION=0` to disable it - or tune it with `BURN_IN_SHIFT_PIXELS` and `BURN_IN_SHIFT_INTERVAL` (seconds). TICI/TIZI and MICI use direct shifting - by default; setting `WHITE_LUMINANCE_CAP` below `1.0` enables the optional luminance cap and offscreen presentation path. +* burn-in prevention shifts the UI by 2 pixels every 3 minutes on device; the shift is blended over 1 second so it is not visible as a jump. + Set `BURN_IN_PREVENTION=0` to disable it or tune it with `BURN_IN_SHIFT_PIXELS`, `BURN_IN_SHIFT_INTERVAL`, and + `BURN_IN_SHIFT_TRANSITION_SECONDS`. TICI/TIZI shift the completed frame through the offscreen presentation path; + MICI uses direct shifting by default. Setting `WHITE_LUMINANCE_CAP` below `1.0` also enables the offscreen presentation path. * set `MICI_FORCE_RENDER_TEXTURE=1` to force the C4 UI through the offscreen presentation path for diagnostics * set `GRID=50` to show a 50-pixel alignment grid overlay * set `MAGIC_DEBUG=1` to show every dropped frames (only on device) diff --git a/system/ui/lib/application.py b/system/ui/lib/application.py index 997172a21..7454de376 100644 --- a/system/ui/lib/application.py +++ b/system/ui/lib/application.py @@ -42,6 +42,10 @@ MICI_FORCE_RENDER_TEXTURE = os.getenv("MICI_FORCE_RENDER_TEXTURE", "0") == "1" BURN_IN_PREVENTION = os.getenv("BURN_IN_PREVENTION", "0" if PC else "1") == "1" BURN_IN_SHIFT_INTERVAL = max(1.0, float(os.getenv("BURN_IN_SHIFT_INTERVAL", "180"))) BURN_IN_SHIFT_PIXELS = max(0, int(os.getenv("BURN_IN_SHIFT_PIXELS", "2"))) +BURN_IN_SHIFT_TRANSITION_SECONDS = min( + BURN_IN_SHIFT_INTERVAL, + max(0.1, float(os.getenv("BURN_IN_SHIFT_TRANSITION_SECONDS", "1"))), +) WHITE_LUMINANCE_CAP = min(1.0, max(0.0, float(os.getenv( "WHITE_LUMINANCE_CAP", "1.0" )))) @@ -616,8 +620,11 @@ class GuiApplication: self._render_texture_width = max(1, int(round(self._scaled_width * self._pixel_scale_x))) self._render_texture_height = max(1, int(round(self._scaled_height * self._pixel_scale_y))) + # Keep raybig burn-in movement in final-frame composition. Translating the live EGL + # camera/widget pass can corrupt the camera presentation instead of shifting the UI. needs_render_texture = ((self._scale != 1.0 and not PC) or BURN_IN_MODE or RECORD or MICI_FORCE_RENDER_TEXTURE or + (BURN_IN_PREVENTION and DEVICE_TYPE != "mici") or WHITE_LUMINANCE_CAP < 1.0) if PC and self._scale != 1.0: rl.set_mouse_scale(1 / self._scale, 1 / self._scale) @@ -1138,13 +1145,25 @@ class GuiApplication: except KeyboardInterrupt: pass - def _burn_in_shift(self, now: float | None = None) -> tuple[int, int]: + def _burn_in_shift(self, now: float | None = None) -> tuple[float, float]: if not BURN_IN_PREVENTION or BURN_IN_SHIFT_PIXELS == 0: - return 0, 0 + return 0.0, 0.0 elapsed = (time.monotonic() if now is None else now) - self._burn_in_start_time - pattern_index = int(max(0.0, elapsed) // BURN_IN_SHIFT_INTERVAL) % len(BURN_IN_SHIFT_PATTERN) - x, y = BURN_IN_SHIFT_PATTERN[pattern_index] + elapsed = max(0.0, elapsed) + pattern_count = len(BURN_IN_SHIFT_PATTERN) + cycle_elapsed = elapsed % (BURN_IN_SHIFT_INTERVAL * pattern_count) + pattern_index = int(cycle_elapsed // BURN_IN_SHIFT_INTERVAL) + segment_elapsed = cycle_elapsed - pattern_index * BURN_IN_SHIFT_INTERVAL + + # Blend into the next position at the end of each interval. This keeps the + # burn-in protection active without teleporting the entire UI by two pixels. + transition_start = BURN_IN_SHIFT_INTERVAL - BURN_IN_SHIFT_TRANSITION_SECONDS + transition = min(1.0, max(0.0, (segment_elapsed - transition_start) / BURN_IN_SHIFT_TRANSITION_SECONDS)) + start_x, start_y = BURN_IN_SHIFT_PATTERN[pattern_index] + end_x, end_y = BURN_IN_SHIFT_PATTERN[(pattern_index + 1) % pattern_count] + x = start_x + (end_x - start_x) * transition + y = start_y + (end_y - start_y) * transition return x * BURN_IN_SHIFT_PIXELS, y * BURN_IN_SHIFT_PIXELS def font(self, font_weight: FontWeight = FontWeight.NORMAL) -> rl.Font: diff --git a/system/ui/lib/tests/test_application.py b/system/ui/lib/tests/test_application.py new file mode 100644 index 000000000..8505497eb --- /dev/null +++ b/system/ui/lib/tests/test_application.py @@ -0,0 +1,16 @@ +from openpilot.system.ui.lib import application + + +def test_burn_in_shift_transitions_between_positions(monkeypatch): + app = object.__new__(application.GuiApplication) + app._burn_in_start_time = 100.0 + + monkeypatch.setattr(application, "BURN_IN_PREVENTION", True) + monkeypatch.setattr(application, "BURN_IN_SHIFT_INTERVAL", 10.0) + monkeypatch.setattr(application, "BURN_IN_SHIFT_PIXELS", 2) + monkeypatch.setattr(application, "BURN_IN_SHIFT_TRANSITION_SECONDS", 2.0) + + assert app._burn_in_shift(108.0) == (0.0, 0.0) + midpoint = app._burn_in_shift(109.0) + assert midpoint == (-1.0, 0.0) + assert app._burn_in_shift(110.0) == (-2.0, 0.0) diff --git a/tinygrad_repo/test/unit/test_usb_gpu_firmware.py b/tinygrad_repo/test/unit/test_usb_gpu_firmware.py deleted file mode 100644 index 09353b8bf..000000000 --- a/tinygrad_repo/test/unit/test_usb_gpu_firmware.py +++ /dev/null @@ -1,37 +0,0 @@ -import pytest - -from tinygrad.runtime.support import system - - -class FakeUSB: - def __init__(self, product: str, is_custom: bool): - self.product, self.is_custom = product, is_custom - - -def patch_usb_dependencies(monkeypatch, usb): - calls = [] - monkeypatch.setattr(system.System, "flock_acquire", lambda _: object()) - monkeypatch.setattr(system, "USB3", lambda *args, **kwargs: calls.append((args, kwargs)) or usb) - return calls - - -def test_usb_gpu_rejects_legacy_bridge_firmware(monkeypatch): - calls = patch_usb_dependencies(monkeypatch, FakeUSB("USB 3.2 PCIe TinyEnclosure", False)) - - with pytest.raises(RuntimeError, match="unsupported legacy USB GPU firmware"): - system.USBPCIDevice("AM", object(), "usb:3-6") - - assert calls[0][1] == {"use_bot": True} - - -def test_usb_gpu_accepts_custom_bridge_firmware(monkeypatch): - usb = FakeUSB("custom ASM2464PD", True) - calls = patch_usb_dependencies(monkeypatch, usb) - controller = object() - monkeypatch.setattr(system, "CustomASM24Controller", lambda candidate: controller if candidate is usb else None) - monkeypatch.setattr(system.System, "pci_setup_usb_bars", lambda *args, **kwargs: {2: (0, 1)}) - - device = system.USBPCIDevice("AM", object(), "usb:3-6") - - assert calls[0][1] == {"use_bot": True} - assert device.usb is controller diff --git a/tinygrad_repo/tinygrad/runtime/ops_amd.py b/tinygrad_repo/tinygrad/runtime/ops_amd.py index 854cb34dd..7404a1fca 100644 --- a/tinygrad_repo/tinygrad/runtime/ops_amd.py +++ b/tinygrad_repo/tinygrad/runtime/ops_amd.py @@ -1,6 +1,6 @@ from __future__ import annotations from typing import cast -import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections, atexit, time +import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections, atexit assert sys.platform != 'win32' from dataclasses import dataclass from tinygrad.runtime.support.hcq import HCQCompiled, HCQAllocator, HCQBuffer, HWQueue, CLikeArgsState, HCQSignal, HCQProgram, FileIOInterface @@ -909,21 +909,15 @@ class PCIIface(PCIIfaceBase): class USBIface(PCIIface): def __init__(self, dev, dev_id): # pylint: disable=super-init-not-called - deadline, visible = time.monotonic() + 5.0, [] - while dev_id >= len(visible) and time.monotonic() < deadline: - visible = hcq_filter_visible_devices(USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001), "AMD") - if dev_id >= len(visible): time.sleep(0.1) - if dev_id >= len(visible): + if dev_id >= len(visible:=hcq_filter_visible_devices(USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001), "AMD")): raise RuntimeError(f"AMD:{dev_id} does not exist ({pluralize('device', len(visible))} available)") self.dev, self.pci_dev, self.vram_bar, self.count = dev, USBPCIDevice("AM", *visible[dev_id]), 0, len(visible) self.dev_impl = AMDev(self.pci_dev) self._compute_props() - self.pci_dev.usb._pci_cacheable += [self.pci_dev.bar_info(2)] # doorbell region is cacheable # special regions self.copy_bufs = [self._dma_region(ctrl_addr=0xf000, sys_addr=0x200000, size=0x80000)] - sys_next_off = 0x200 if self.pci_dev.usb.usb.is_custom else 0x800 - self.sys_buf, self.sys_next_off = self._dma_region(ctrl_addr=0xa000, sys_addr=0x820000, size=0x1000), sys_next_off + self.sys_buf, self.sys_next_off = self._dma_region(ctrl_addr=0xa000, sys_addr=0x820000, size=0x1000), 0x200 self.cq_buf = self._dma_region(ctrl_addr=0xb800, sys_addr=0x822000, size=0x1000) def _dma_region(self, ctrl_addr, sys_addr, size): @@ -932,18 +926,13 @@ class USBIface(PCIIface): def alloc(self, size:int, host=False, uncached=False, cpu_access=False, contiguous=False, force_devmem=False, **kwargs) -> HCQBuffer: # usb allocates uncached and cpu_access in vram. vram writes are faster than sram writes - if (host or (not self.pci_dev.usb.usb.is_custom and uncached and cpu_access)) and self.sys_next_off + size < self.sys_buf.size: + if host and self.sys_next_off + size < self.sys_buf.size: self.sys_next_off += size return self.sys_buf.offset(self.sys_next_off - size, size) # force devmem return super().alloc(size, host=False, uncached=uncached, cpu_access=cpu_access, contiguous=contiguous, force_devmem=True, **kwargs) - def create_queue(self, queue_type, ring, gart, rptr, wptr, eop_buffer=None, cwsr_buffer=None, ctl_stack_size=0, ctx_save_restore_size=0, - xcc_id=0, idx=0): - if queue_type == kfd.KFD_IOC_QUEUE_TYPE_COMPUTE: self.pci_dev.usb._pci_cacheable += [(ring.cpu_view().addr, ring.size)] - return super().create_queue(queue_type, ring, gart, rptr, wptr, eop_buffer, cwsr_buffer, ctl_stack_size, ctx_save_restore_size, xcc_id, idx) - def sleep(self, timeout): pass def _mock(iface, name=None): return type(name or f"MOCK{iface.__name__}", (iface,), {}) diff --git a/tinygrad_repo/tinygrad/runtime/support/system.py b/tinygrad_repo/tinygrad/runtime/support/system.py index a7c8a11f5..44e560618 100644 --- a/tinygrad_repo/tinygrad/runtime/support/system.py +++ b/tinygrad_repo/tinygrad/runtime/support/system.py @@ -225,10 +225,8 @@ class USBPCIDevice(PCIDevice): def __init__(self, devpref:str, dev, pcibus): self.pcibus, self.peer_group = pcibus, f"USBPCIDevice_{pcibus}" self.lock_fd = System.flock_acquire(f"{devpref.lower()}_{pcibus.lower()}.lock") - usb = USB3(dev, 0x81, 0x83, 0x02, 0x04, use_bot=True) + usb = USB3(dev) if DEBUG >= 1: print(f"am {self.pcibus}: product string: {usb.product!r}") - if not usb.is_custom: - raise RuntimeError(f"unsupported legacy USB GPU firmware ({usb.product!r}); flash the current tinygrad ASM2464PD custom firmware") self.usb: CustomASM24Controller = CustomASM24Controller(usb) self._bar_info = System.pci_setup_usb_bars(self.usb, gpu_bus=4, mem_base=0x10000000, pref_mem_base=(32 << 30)) self.sram = BumpAllocator(size=0x80000, wrap=False) # asm24 controller sram diff --git a/tinygrad_repo/tinygrad/runtime/support/usb.py b/tinygrad_repo/tinygrad/runtime/support/usb.py index 37cf6815a..b7e8b8186 100644 --- a/tinygrad_repo/tinygrad/runtime/support/usb.py +++ b/tinygrad_repo/tinygrad/runtime/support/usb.py @@ -1,7 +1,6 @@ -import ctypes, struct, dataclasses, array, itertools, time, functools -from typing import Sequence +import ctypes, struct, time, functools, itertools from tinygrad.runtime.autogen import libusb -from tinygrad.helpers import DEBUG, DEV, to_mv, round_up, OSX, getenv, ceildiv +from tinygrad.helpers import DEBUG, DEV, to_mv, round_up, ceildiv from tinygrad.runtime.support.hcq import MMIOInterface from tinygrad.runtime.support import c @@ -22,6 +21,7 @@ class USB3: return ctx @classmethod + @functools.cache def list_devices(cls, vendor:int, dev:int) -> list[tuple[c.POINTER[libusb.struct_libusb_device], str]]: ret = [] for i in range(checked(libusb.libusb_get_device_list)(cls.ctx(), devs:=ctypes.POINTER(ctypes.POINTER(libusb.struct_libusb_device))())): @@ -31,31 +31,12 @@ class USB3: libusb.libusb_free_device_list(devs, 1) return ret - @classmethod - def reopen_device(cls, vendor:int, product:int, bus_number:int, timeout:float=10.0): - deadline = time.monotonic() + timeout - while True: - for dev, _ in cls.list_devices(vendor, product): - handle = ctypes.POINTER(libusb.struct_libusb_device_handle)() - rc = libusb.libusb_open(dev, ctypes.byref(handle)) if libusb.libusb_get_bus_number(dev) == bus_number else libusb.LIBUSB_ERROR_NOT_FOUND - libusb.libusb_unref_device(dev) - if rc == 0: - return handle - if time.monotonic() >= deadline: - raise RuntimeError(f"device {vendor:04x}:{product:04x} did not reappear after reset") - time.sleep(0.1) + def __init__(self, dev:c.POINTER[libusb.struct_libusb_device], *args, **kwargs): + self._tags, self._transferred = itertools.count(1), ctypes.c_int(0) + self._bulk_buf, self._bulk_mv = alloc_cbuffer(4 << 20) + self._ctrl_buf, self._ctrl_mv = alloc_cbuffer(0x1000) - def __init__(self, dev:c.POINTER[libusb.struct_libusb_device], ep_data_in:int, ep_stat_in:int, ep_data_out:int, ep_cmd_out:int, - max_streams:int=31, use_bot=False): - self.ep_data_in, self.ep_stat_in, self.ep_data_out, self.ep_cmd_out = ep_data_in, ep_stat_in, ep_data_out, ep_cmd_out - self.max_streams, self.use_bot = max_streams, use_bot - self._transferred = ctypes.c_int(0) - self._bulk_in_buf, self._bulk_in_mv = alloc_cbuffer(4 << 20) - self._bulk_out_buf, self._bulk_out_mv = alloc_cbuffer(4 << 20) - - bus_number = libusb.libusb_get_bus_number(dev) - self.handle = c.init_c_var(ctypes.POINTER(libusb.struct_libusb_device_handle), lambda x: checked(libusb.libusb_open)(dev, x)) - libusb.libusb_unref_device(dev) + self.handle = c.init_c_var(c.POINTER[libusb.struct_libusb_device_handle], lambda x: checked(libusb.libusb_open)(dev, x)) # Read product string descriptor _buf = (ctypes.c_ubyte * 256)() @@ -63,210 +44,83 @@ class USB3: checked(libusb.libusb_get_device_descriptor)(libusb.libusb_get_device(self.handle), ctypes.byref(_desc)) _ret = checked(libusb.libusb_get_string_descriptor_ascii)(self.handle, _desc.iProduct, _buf, 256) self.product = bytes(_buf[:_ret]).decode("ascii", errors="replace") - self.is_custom = self.product.startswith("custom") - if self.is_custom: self.use_bot = use_bot = True + assert self.product.startswith("custom") or self.product.startswith("AS2462") # Detach kernel driver if needed if checked(libusb.libusb_kernel_driver_active)(self.handle, 0): checked(libusb.libusb_detach_kernel_driver)(self.handle, 0) - reset_rc = libusb.libusb_reset_device(self.handle) - if reset_rc in (libusb.LIBUSB_ERROR_NO_DEVICE, libusb.LIBUSB_ERROR_NOT_FOUND): - libusb.libusb_close(self.handle) - self.handle = self.reopen_device(_desc.idVendor, _desc.idProduct, bus_number) - if checked(libusb.libusb_kernel_driver_active)(self.handle, 0): - checked(libusb.libusb_detach_kernel_driver)(self.handle, 0) - elif reset_rc < 0: - raise RuntimeError(f"libusb_reset_device: {ctypes.string_at(libusb.libusb_strerror(reset_rc)).decode()}") + checked(libusb.libusb_reset_device)(self.handle) # Set configuration and claim interface checked(libusb.libusb_set_configuration)(self.handle, 1) checked(libusb.libusb_claim_interface)(self.handle, 0) + checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 0) - if use_bot: - checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 0) - self._tag = 0 - else: - checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 1) + def control_write(self, request:int, value:int=0, index:int=0, data:bytes=b'', timeout:int=1000): + assert len(data) <= len(self._ctrl_mv) + self._ctrl_mv[:len(data)] = data + assert checked(libusb.libusb_control_transfer)(self.handle, 0x40, request, value, index, self._ctrl_buf, len(data), timeout) == len(data) - # Clear any stalled endpoints - all_eps = (self.ep_data_out, self.ep_data_in, self.ep_stat_in, self.ep_cmd_out) - for ep in all_eps: checked(libusb.libusb_clear_halt)(self.handle, ep) + def control_read(self, request:int, length:int, value:int=0, index:int=0, timeout:int=1000) -> memoryview: + assert length <= len(self._ctrl_mv) + assert checked(libusb.libusb_control_transfer)(self.handle, 0xC0, request, value, index, self._ctrl_buf, length, timeout) == length + return self._ctrl_mv[:length] - # Allocate streams - stream_eps = (ctypes.c_uint8 * 3)(self.ep_data_out, self.ep_data_in, self.ep_stat_in) - checked(libusb.libusb_alloc_streams)(self.handle, self.max_streams * len(stream_eps), stream_eps, len(stream_eps)) + def bulk_write(self, payload:bytes, timeout:int=1000): + if len(payload) > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(len(payload)) + self._bulk_mv[:len(payload)] = payload + checked(libusb.libusb_bulk_transfer, "bulk OUT 0x02 failed") \ + (self.handle, 0x02, self._bulk_buf, len(payload), self._transferred, timeout) + assert self._transferred.value == len(payload), f"bulk OUT short write: {self._transferred.value}/{len(payload)} bytes" - # Base cmd - cmd_template = bytes([0x01, 0x00, 0x00, 0x01, *([0] * 12), 0xE4, 0x24, 0x00, 0xB2, 0x1A, 0x00, 0x00, 0x00, *([0] * 8)]) + def bulk_read(self, length:int, timeout:int=1000) -> memoryview: + if length > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(length) + checked(libusb.libusb_bulk_transfer, "bulk IN 0x81 failed")(self.handle, 0x81, self._bulk_buf, length, self._transferred, timeout) + return self._bulk_mv[:self._transferred.value] - # Init pools - self.tr = {ep: [libusb.libusb_alloc_transfer(0) for _ in range(self.max_streams)] for ep in all_eps} - - self.buf_cmd = [(ctypes.c_uint8 * len(cmd_template))(*cmd_template) for _ in range(self.max_streams)] - self.buf_stat = [(ctypes.c_uint8 * 64)() for _ in range(self.max_streams)] - self.buf_data_in = [(ctypes.c_uint8 * 0x1000)() for _ in range(self.max_streams)] - self.buf_data_out = [(ctypes.c_uint8 * 0x80000)() for _ in range(self.max_streams)] - self.buf_data_out_mvs = [to_mv(ctypes.addressof(self.buf_data_out[i]), 0x80000) for i in range(self.max_streams)] - - for slot in range(self.max_streams): struct.pack_into(">B", self.buf_cmd[slot], 3, slot + 1) - - def _prep_transfer(self, tr, ep, stream_id, buf, length): - tr.contents.dev_handle, tr.contents.endpoint, tr.contents.length, tr.contents.buffer = self.handle, ep, length, buf - tr.contents.status, tr.contents.flags, tr.contents.timeout, tr.contents.num_iso_packets = 0xff, 0, 1000, 0 - tr.contents.type = (libusb.LIBUSB_TRANSFER_TYPE_BULK_STREAM if stream_id is not None else libusb.LIBUSB_TRANSFER_TYPE_BULK) - if stream_id is not None: libusb.libusb_transfer_set_stream_id(tr, stream_id) - return tr - - def _submit_and_wait(self, cmds): - for tr in cmds: checked(libusb.libusb_submit_transfer)(tr) - - running = len(cmds) - while running: - checked(libusb.libusb_handle_events)(USB3.ctx()) - running = len(cmds) - for tr in cmds: - if tr.contents.status == libusb.LIBUSB_TRANSFER_COMPLETED: running -= 1 - elif tr.contents.status != 0xFF: raise RuntimeError(f"EP 0x{tr.contents.endpoint:02X} error: {tr.contents.status}") - - def _bulk_out(self, ep: int, payload: bytes, timeout: int = 1000): - if len(payload) > len(self._bulk_out_mv): self._bulk_out_buf, self._bulk_out_mv = alloc_cbuffer(len(payload)) - self._bulk_out_mv[:len(payload)] = payload - checked(libusb.libusb_bulk_transfer, f"bulk OUT 0x{ep:02X} failed")(self.handle, ep, self._bulk_out_buf, len(payload), self._transferred, timeout) - assert self._transferred.value == len(payload), f"bulk OUT short write on 0x{ep:02X}: {self._transferred.value}/{len(payload)} bytes" - - def _bulk_in(self, ep: int, length: int, timeout: int = 1000) -> memoryview: - if length > len(self._bulk_in_mv): self._bulk_in_buf, self._bulk_in_mv = alloc_cbuffer(length) - checked(libusb.libusb_bulk_transfer, f"bulk IN 0x{ep:02X} failed")(self.handle, ep, self._bulk_in_buf, length, self._transferred, timeout) - return self._bulk_in_mv[:self._transferred.value] - - def send_batch(self, cdbs:list[bytes], idata:list[int]|None=None, odata:list[bytes|None]|None=None) -> list[bytes|None]: - idata, odata = idata or [0] * len(cdbs), odata or [None] * len(cdbs) - results:list[bytes|None] = [] - tr_window, op_window = [], [] - - for idx, (cdb, rlen, send_data) in enumerate(zip(cdbs, idata, odata)): - if self.use_bot: - dir_in = rlen > 0 - data_len = rlen if dir_in else (len(send_data) if send_data is not None else 0) - assert not (rlen > 0 and send_data is not None), "BOT mode only supports either read or write per command" - - # CBW - self._tag += 1 - flags = 0x80 if dir_in else 0x00 - cbw = struct.pack(" len(self.buf_data_in[slot]): self.buf_data_in[slot] = (ctypes.c_uint8 * round_up(rlen, 0x1000))() - tr_window.append(self._prep_transfer(self.tr[self.ep_data_in][slot], self.ep_data_in, stream, self.buf_data_in[slot], rlen)) - - if send_data is not None: - if len(send_data) > len(self.buf_data_out[slot]): - self.buf_data_out[slot] = (ctypes.c_uint8 * len(send_data))() - self.buf_data_out_mvs[slot] = to_mv(ctypes.addressof(self.buf_data_out[slot]), len(send_data)) - - self.buf_data_out_mvs[slot][:len(send_data)] = bytes(send_data) - tr_window.append(self._prep_transfer(self.tr[self.ep_data_out][slot], self.ep_data_out, stream, self.buf_data_out[slot], len(send_data))) - - op_window.append((idx, slot, rlen)) - if (idx + 1 == len(cdbs)) or len(op_window) >= self.max_streams: - self._submit_and_wait(tr_window) - for idx, slot, rlen in op_window: results.append(bytes(self.buf_data_in[slot][:rlen]) if rlen else None) - tr_window = [] - - return results - -@dataclasses.dataclass(frozen=True) -class WriteOp: addr:int; data:bytes; ignore_cache:bool=True # noqa: E702 - -@dataclasses.dataclass(frozen=True) -class ReadOp: addr:int; size:int # noqa: E702 - -@dataclasses.dataclass(frozen=True) -class ScsiWriteOp: data:bytes; lba:int=0 # noqa: E702 + # NOTE: keep it for flash.py + def send_batch(self, cdbs:list[bytes], odata:list[bytes|None]|None=None): + for cdb, data in zip(cdbs, odata or [None] * len(cdbs)): + self.bulk_write(struct.pack("> 32, value) - ret = libusb.libusb_control_transfer(self.usb.handle, 0x40, 0xF0, fmt_type | (byte_en << 8), mode & 0x03, self._f0_out_buf, 12, 5000) - assert ret == 12, f"F0 OUT failed: {ret}" + self.usb.control_write(0xF0, fmt_type | (byte_en << 8), mode & 0x03, struct.pack('> 32, value), 5000) def _f0_in(self) -> tuple[int, int, int]: - ret = libusb.libusb_control_transfer(self.usb.handle, 0xC0, 0xF0, 0, 0, self._f0_in_buf, 8, 5000) - assert ret == 8, f"F0 IN failed: {ret}" - return struct.unpack_from('> 5) & 0x7, self._f0_in_buf[7] - - def _is_pci_cacheable(self, addr:int) -> bool: return any(x <= addr <= x + sz for x, sz in self._pci_cacheable) + data = self.usb.control_read(0xF0, 8, timeout=5000) + return struct.unpack_from('> 5) & 0x7, data[7] def pcie_request(self, fmt_type:int, address:int, value:int|None=None, size:int=4, cnt:int=10): - if fmt_type == 0x60 and size == 4 and self._is_pci_cacheable(address) and self._pci_cache.get(address) == value: return assert size > 0 and size <= 4, f"Invalid size {size}" if DEBUG >= 5: print("pcie_request", hex(fmt_type), hex(address), value, size) offset = address & 0x3 byte_en = ((1 << size) - 1) << offset - self._pci_cache[address] = value if size == 4 and fmt_type == 0x60 else None - self._f0_out(fmt_type, byte_en, address & ~0x3, (value << (8 * offset)) if value is not None else 0) - # Fast path: memory writes and messages don't return completions (same logic as ASM24Controller). + # Fast path: memory writes and messages don't return completions. if ((fmt_type & 0b11011111) == 0b01000000) or ((fmt_type & 0b10111000) == 0b00110000): return # Read TLPs and config writes: read completion via 0xF0 IN. Retry on error/timeout. data, cpl_status, ret_status = self._f0_in() if ret_status != 0: time.sleep(0.001) # TODO: this sleep is very picky - if cnt > 0: - return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1) + if cnt > 0: return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1) raise RuntimeError(f"TLP error after retries: ret_status={ret_status}, address={address:#x}") if cpl_status: @@ -281,219 +135,69 @@ class CustomASM24Controller: address = (bus << 24) | (dev << 19) | (fn << 16) | (byte_addr & 0xfff) return self.pcie_request(fmt_type, address, value, size) - def pcie_mem_req(self, address:int, value:int|None=None, size:int=4): - return self.pcie_request(0x60 if value is not None else 0x20, address, value, size) - - def pcie_mem_write(self, address:int, values:list[int], size:int): + def pcie_mem_write(self, address:int, data:bytes): """Streaming PCIe memory write via 0xF0 mode 1 + bulk OUT. Data is little-endian dwords on the wire.""" - if not values: return - self._f0_out(0x60, 0x0F, address, len(values), mode=1) - self.usb._bulk_out(0x02, struct.pack(f'<{len(values)}I', *values)) + if not data: return + assert len(data) % 4 == 0, f"pcie_mem_write requires 4-byte aligned size, got {len(data)}" + self._f0_out(0x60, 0x0F, address, len(data) // 4, mode=1) + self.usb.bulk_write(data) - def pcie_mem_read(self, address:int, nbytes:int) -> bytes: + def pcie_mem_read(self, address:int, nbytes:int) -> memoryview: """Streaming PCIe memory read via 0xF0 mode 2 + bulk IN. Returns little-endian bytes.""" assert nbytes % 4 == 0, f"pcie_mem_read requires 4-byte aligned size, got {nbytes}" self._f0_out(0x20, 0x0F, address, nbytes // 4, mode=2) - return self.usb._bulk_in(0x81, nbytes, timeout=30000) + return self.usb.bulk_read(nbytes, timeout=30000) - # === XDATA read/write (0xE4/0xE5 vendor control transfers) === - - def read(self, base_addr:int, length:int, **kwargs) -> bytes: + def read(self, base_addr:int, length:int) -> bytes: """Read from chip XDATA via vendor control IN (bRequest=0xE4). wValue=addr, wLength=size.""" result = b'' for off in range(0, length, 0xFF): chunk = min(0xFF, length - off) - ret = libusb.libusb_control_transfer(self.usb.handle, 0xC0, 0xE4, base_addr + off, 0, self._f0_out_buf, chunk, 1000) - assert ret == chunk, f"read(0x{base_addr + off:04X}, {chunk}) failed: {ret}" - result += bytes(self._f0_out_buf[:ret]) - return result[:length] + result += self.usb.control_read(0xE4, chunk, value=base_addr + off) + return result - def write(self, base_addr:int, data:bytes, **kwargs): + def write(self, base_addr:int, data:bytes): """Write to chip XDATA via vendor control OUT (bRequest=0xE5). wValue=addr, wIndex=val.""" - for off, val in enumerate(data): - checked(libusb.libusb_control_transfer, - f"write(0x{base_addr + off:04X}, 0x{val:02X}) failed")(self.usb.handle, 0x40, 0xE5, base_addr + off, val, None, 0, 1000) + for off, val in enumerate(data): self.usb.control_write(0xE5, value=base_addr + off, index=val) - def scsi_write(self, buf:bytes, lba:int=0): + def scsi_write(self, buf:bytes): """Write to SRAM via 0xF2 vendor command + bulk OUT.""" buf_padded = buf + b'\x00' * (round_up(len(buf), 512) - len(buf)) sectors = len(buf_padded) // 512 - num_slots = round_up(len(buf_padded), 0x4000) // 0x4000 # 16KB per slot - # 0xF2 OUT: wValue=sectors, wIndex=start_slot|(num_slots<<8) + num_slots = ceildiv(len(buf_padded), 0x4000) # 16KB per slot windex = (num_slots & 0xFF) << 8 - checked(libusb.libusb_control_transfer, "F2 setup failed")(self.usb.handle, 0x40, 0xF2, sectors, windex, None, 0, 1000) - self.usb._bulk_out(0x02, buf_padded) + self.usb.control_write(0xF2, value=sectors, index=windex) + self.usb.bulk_write(buf_padded) def scsi_read_arm(self, size:int): windex = (ceildiv(size, 0x4000) & 0xFF) << 8 - checked(libusb.libusb_control_transfer, - "F2 read arm failed")(self.usb.handle, 0x40, 0xF2, (ceildiv(size, 512) & 0x7FFF) | 0x8000, windex, None, 0, 1000) + self.usb.control_write(0xF2, value=(ceildiv(size, 512) & 0x7FFF) | 0x8000, index=windex) - def scsi_read(self, size:int) -> memoryview: return self.usb._bulk_in(0x81, round_up(size, 512), timeout=10000)[:size] - -class ASM24Controller: - def __init__(self, usb:USB3|None=None): - if not usb: - devs = USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001) - assert len(devs), "no ASM24 controller found" - self.usb = USB3(devs[0][0], 0x81, 0x83, 0x02, 0x04, use_bot=bool(getenv("USE_BOT", 0))) - else: self.usb = usb - self._cache: dict[int, int|None] = {} - self._pci_cacheable: list[tuple[int, int]] = [] - self._pci_cache: dict[int, int|None] = {} - - # Init controller. - self.exec_ops([WriteOp(0x54b, b' '), WriteOp(0x54e, b'\x04'), WriteOp(0x5a8, b'\x02'), WriteOp(0x5f8, b'\x04'), - WriteOp(0x7ec, b'\x01\x00\x00\x00'), WriteOp(0xc422, b'\x02'), WriteOp(0x0, b'\x33')]) - - def exec_ops(self, ops:Sequence[WriteOp|ReadOp|ScsiWriteOp]): - cdbs:list[bytes] = [] - idata:list[int] = [] - odata:list[bytes|None] = [] - - def _add_req(cdb:bytes, i:int, o:bytes|None): - nonlocal cdbs, idata, odata - cdbs, idata, odata = cdbs + [cdb], idata + [i], odata + [o] - - for op in ops: - if isinstance(op, WriteOp): - for off, value in enumerate(op.data): - addr = ((op.addr + off) & 0x1FFFF) | 0x500000 - if not op.ignore_cache and self._cache.get(addr) == value: continue - _add_req(struct.pack('>BBBHB', 0xE5, value, addr >> 16, addr & 0xFFFF, 0), 0, None) - self._cache[addr] = value - elif isinstance(op, ReadOp): - assert op.size <= 0xff - addr = (op.addr & 0x1FFFF) | 0x500000 - _add_req(struct.pack('>BBBHB', 0xE4, op.size, addr >> 16, addr & 0xFFFF, 0), op.size, None) - for i in range(op.size): self._cache[addr + i] = None - elif isinstance(op, ScsiWriteOp): - sectors = round_up(len(op.data), 512) // 512 - _add_req(struct.pack('>BBQIBB', 0x8A, 0, op.lba, sectors, 0, 0), 0, op.data+b'\x00'*((sectors*512)-len(op.data))) - - return self.usb.send_batch(cdbs, idata, odata) - - def write(self, base_addr:int, data:bytes, ignore_cache:bool=True): return self.exec_ops([WriteOp(base_addr, data, ignore_cache)]) - - def scsi_write(self, buf:bytes, lba:int=0): - if len(buf) > 0x4000: buf += b'\x00' * (round_up(len(buf), 0x10000) - len(buf)) - - for i in range(0, len(buf), 0x10000): - self.exec_ops([ScsiWriteOp(buf[i:i+0x10000], lba), WriteOp(0x171, b'\xff\xff\xff', ignore_cache=True)]) - self.exec_ops([WriteOp(0xce6e, b'\x00\x00', ignore_cache=True)]) - - if len(buf) > 0x4000: - for i in range(4): self.exec_ops([WriteOp(0xce40 + i, b'\x00', ignore_cache=True)]) - - def read(self, base_addr:int, length:int, stride:int=0xff) -> bytes: - parts = self.exec_ops([ReadOp(base_addr + off, min(stride, length - off)) for off in range(0, length, stride)]) - return b''.join(p or b'' for p in parts)[:length] - - def _is_pci_cacheable(self, addr:int) -> bool: return any(x <= addr <= x + sz for x, sz in self._pci_cacheable) - def pcie_prep_request(self, fmt_type:int, address:int, value:int|None=None, size:int=4) -> list[WriteOp]: - if fmt_type == 0x60 and size == 4 and self._is_pci_cacheable(address) and self._pci_cache.get(address) == value: return [] - - assert fmt_type >> 8 == 0 and size > 0 and size <= 4, f"Invalid fmt_type {fmt_type} or size {size}" - if DEBUG >= 5: print("pcie_request", hex(fmt_type), hex(address), value, size) - - masked_address, offset = address & 0xFFFFFFFC, address & 0x3 - assert size + offset <= 4 and (value is None or value >> (8 * size) == 0) - self._pci_cache[address] = value if size == 4 and fmt_type == 0x60 else None - - return ([WriteOp(0xB220, struct.pack('>I', value << (8 * offset)), ignore_cache=False)] if value is not None else []) + \ - [WriteOp(0xB218, struct.pack('>I', masked_address), ignore_cache=False), WriteOp(0xB21c, struct.pack('>I', address>>32), ignore_cache=False), - WriteOp(0xB217, bytes([((1 << size) - 1) << offset]), ignore_cache=False), WriteOp(0xB210, bytes([fmt_type]), ignore_cache=False), - WriteOp(0xB254, b"\x0f", ignore_cache=True), WriteOp(0xB296, b"\x04", ignore_cache=True)] - - def pcie_request(self, fmt_type, address, value=None, size=4, cnt=10): - self.exec_ops(self.pcie_prep_request(fmt_type, address, value, size)) - - # Fast path for write requests - if ((fmt_type & 0b11011111) == 0b01000000) or ((fmt_type & 0b10111000) == 0b00110000): return - - while (stat:=self.read(0xB296, 1)[0]) & 2 == 0: - if stat & 1: - self.write(0xB296, bytes([0x01])) - if cnt > 0: return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1) - assert stat == 2, f"stat read 2 was {stat}" - - # Retrieve completion data from Link Status (0xB22A, 0xB22B) - b284 = self.read(0xB284, 1)[0] - completion = struct.unpack('>H', self.read(0xB22A, 2)) - - # Validate completion status based on PCIe request typ - # Completion TLPs for configuration requests always have a byte count of 4. - assert completion[0] & 0xfff == (4 if (fmt_type & 0xbe == 0x04) else size) - - # Extract completion status field - status = (completion[0] >> 13) & 0x7 - - # Handle completion errors or inconsistencies - if status or ((fmt_type & 0xbe == 0x04) and (((value is None) and (not (b284 & 0x01))) or ((value is not None) and (b284 & 0x01)))): - status_map = {0b001: f"Unsupported Request: invalid address/function (target might not be reachable): {address:#x}", - 0b100: "Completer Abort: abort due to internal error", 0b010: "Configuration Request Retry Status: configuration space busy"} - raise RuntimeError(f"TLP status: {status_map.get(status, 'Reserved (0b{:03b})'.format(status))}") - - if value is None: return (struct.unpack('>I', self.read(0xB220, 4))[0] >> (8 * (address & 0x3))) & ((1 << (8 * size)) - 1) - - def pcie_cfg_req(self, byte_addr, bus=1, dev=0, fn=0, value=None, size=4): - assert byte_addr >> 12 == 0 and bus >> 8 == 0 and dev >> 5 == 0 and fn >> 3 == 0, f"Invalid byte_addr {byte_addr}, bus {bus}, dev {dev}, fn {fn}" - - fmt_type = (0x44 if value is not None else 0x4) | int(bus > 0) - address = (bus << 24) | (dev << 19) | (fn << 16) | (byte_addr & 0xfff) - return self.pcie_request(fmt_type, address, value, size) - - def pcie_mem_req(self, address, value=None, size=4): return self.pcie_request(0x60 if value is not None else 0x20, address, value, size) - - def pcie_mem_write(self, address, values, size): - ops = [self.pcie_prep_request(0x60, address + i * size, value, size) for i, value in enumerate(values)] - - # Send in batches of 4 for OSX and 16 for Linux (benchmarked values) - for i in range(0, len(ops), bs:=(4 if OSX else 16)): self.exec_ops(list(itertools.chain.from_iterable(ops[i:i+bs]))) + def scsi_read(self, size:int) -> memoryview: return self.usb.bulk_read(round_up(size, 512), timeout=10000)[:size] class USBMMIOInterface(MMIOInterface): def __init__(self, usb, addr, size, fmt, pcimem=True): # pylint: disable=super-init-not-called - self.usb, self.addr, self.nbytes, self.fmt, self.pcimem, self.el_sz = usb, addr, size, fmt, pcimem, struct.calcsize(fmt) + self.usb, self.addr, self.nbytes, self.fmt, self.el_sz, self.pcimem = usb, addr, size, fmt, struct.calcsize(fmt), pcimem - def __getitem__(self, index): return self._access_items(index) - def __setitem__(self, index, val): self._access_items(index, val) + def _off_from_index(self, index): + if isinstance(index, slice): return ((index.start or 0) * self.el_sz, ((index.stop or len(self))-(index.start or 0)) * self.el_sz) + return (index * self.el_sz, self.el_sz) - def _access_items(self, index, val=None): - if isinstance(index, slice): return self._acc((index.start or 0) * self.el_sz, ((index.stop or len(self))-(index.start or 0)) * self.el_sz, val) - return self._acc_one(index * self.el_sz, self.el_sz, val) if self.pcimem else self._acc(index * self.el_sz, self.el_sz, val) + def __getitem__(self, index): + off, sz = self._off_from_index(index) + if self.pcimem: + assert sz % 4 == 0 and off % 4 == 0, f"pcie_mem_read requires 4-byte aligned access, got off={off}, sz={sz}" + data = self.usb.pcie_mem_read(self.addr + off, sz) + else: data = self.usb.scsi_read(sz) if self.addr == 0xf000 else self.usb.read(self.addr + off, sz) + return int.from_bytes(data, "little") if sz == self.el_sz else data + + def __setitem__(self, index, data): + off, _ = self._off_from_index(index) + data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data) + if not self.pcimem: self.usb.scsi_write(data) if self.addr == 0xf000 else self.usb.write(self.addr + off, data) + else: self.usb.pcie_mem_write(self.addr+off, data) def view(self, offset:int=0, size:int|None=None, fmt=None): - return USBMMIOInterface(self.usb, self.addr+offset, size or (self.nbytes - offset), fmt=fmt or self.fmt, pcimem=self.pcimem) - - def _acc_size(self, sz): return next(x for x in [('I', 4), ('H', 2), ('B', 1)] if sz % x[1] == 0) - - def _acc_one(self, off, sz, val=None): - upper = 0 if sz < 8 else self.usb.pcie_mem_req(self.addr + off + 4, val if val is None else (val >> 32), 4) - lower = self.usb.pcie_mem_req(self.addr + off, val if val is None else val & 0xffffffff, min(sz, 4)) - if val is None: return lower | (upper << 32) - - def _acc(self, off, sz, data=None): - if data is None: # read op - if not self.pcimem: - if self.addr == 0xf000 and hasattr(self.usb, 'scsi_read'): return self.usb.scsi_read(sz) - return int.from_bytes(self.usb.read(self.addr + off, sz), "little") if sz == self.el_sz else self.usb.read(self.addr + off, sz) - - # Fast path: streaming PCIe read if controller supports it - if hasattr(self.usb, 'pcie_mem_read') and sz >= 4 and sz % 4 == 0: - return self.usb.pcie_mem_read(self.addr + off, sz) - - acc, acc_size = self._acc_size(sz) - return bytes(array.array(acc, [self._acc_one(off + i * acc_size, acc_size) for i in range(sz // acc_size)])) - - # write op - data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data) - - if not self.pcimem: - # Fast path for writing into buffer 0xf000 - use_cache = 0xa800 <= self.addr <= 0xb000 - return self.usb.scsi_write(bytes(data)) if self.addr == 0xf000 else self.usb.write(self.addr + off, bytes(data), ignore_cache=not use_cache) - - _, acc_sz = self._acc_size(len(data) * struct.calcsize(self.fmt)) - self.usb.pcie_mem_write(self.addr+off, [int.from_bytes(data[i:i+acc_sz], "little") for i in range(0, len(data), acc_sz)], acc_sz) + return USBMMIOInterface(self.usb, self.addr+offset, self.nbytes-offset if size is None else size, fmt=fmt or self.fmt, pcimem=self.pcimem) if DEV.interface.startswith("MOCK"): from test.mockgpu.usb import MockUSB3 as USB3 # type: ignore # noqa: F811