mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-08-20 07:43:48 +08:00
The biggest, the largest
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]))],
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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}')
|
||||
|
||||
@@ -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)},
|
||||
}
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
+4
-3
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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,), {})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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("<IIIBBB", 0x43425355, self._tag, data_len, flags, 0, len(cdb)) + cdb + b"\x00" * (16 - len(cdb))
|
||||
self._bulk_out(self.ep_data_out, cbw)
|
||||
|
||||
# DAT
|
||||
if dir_in:
|
||||
results.append(bytes(self._bulk_in(self.ep_data_in, rlen)))
|
||||
else:
|
||||
if send_data is not None:
|
||||
self._bulk_out(self.ep_data_out, send_data)
|
||||
results.append(None)
|
||||
|
||||
# CSW
|
||||
sig, rtag, residue, status = struct.unpack("<IIIB", self._bulk_in(self.ep_data_in, 13, timeout=2000))
|
||||
assert sig == 0x53425355, f"Bad CSW signature 0x{sig:08X}, expected 0x53425355"
|
||||
assert rtag == self._tag, f"CSW tag mismatch: got {rtag}, expected {self._tag}"
|
||||
assert status == 0, f"SCSI command failed, CSW status=0x{status:02X}, residue={residue}"
|
||||
else:
|
||||
# allocate slot and stream. stream is 1-based
|
||||
slot, stream = idx % self.max_streams, (idx % self.max_streams) + 1
|
||||
|
||||
# build cmd packet
|
||||
self.buf_cmd[slot][16:16+len(cdb)] = list(cdb)
|
||||
|
||||
# cmd + stat transfers
|
||||
tr_window.append(self._prep_transfer(self.tr[self.ep_cmd_out][slot], self.ep_cmd_out, None, self.buf_cmd[slot], len(self.buf_cmd[slot])))
|
||||
tr_window.append(self._prep_transfer(self.tr[self.ep_stat_in][slot], self.ep_stat_in, stream, self.buf_stat[slot], 64))
|
||||
|
||||
if rlen:
|
||||
if rlen > 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("<IIIBBB16s", 0x43425355, tag:=next(self._tags), len(data) if data is not None else 0, 0, 0, len(cdb), cdb))
|
||||
if data is not None: self.bulk_write(data)
|
||||
sig, rtag, _, status = struct.unpack("<IIIB", self.bulk_read(13, timeout=2000))
|
||||
assert (sig, rtag, status) == (0x53425355, tag, 0)
|
||||
|
||||
class CustomASM24Controller:
|
||||
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=True)
|
||||
else: self.usb = usb
|
||||
self._pci_cacheable: list[tuple[int, int]] = []
|
||||
self._pci_cache: dict[int, int|None] = {}
|
||||
def __init__(self, usb:USB3):
|
||||
self.usb = usb
|
||||
|
||||
self._f0_out_buf, self._f0_out_mv = alloc_cbuffer(0x1000) # for f0 and e4, allocate big enough for e4
|
||||
self._f0_in_buf, _ = alloc_cbuffer(8)
|
||||
|
||||
if (ltssm := self.read(0xB450, 1)[0]) != 0x78:
|
||||
self.set_pcie_power(True)
|
||||
grace_period = time.monotonic() + 5.0
|
||||
while time.monotonic() < grace_period and (ltssm := self.read(0xB450, 1)[0]) != 0x78:
|
||||
time.sleep(0.1)
|
||||
# Custom firmware now boots with PCIe off. Power it on before probing the link.
|
||||
ltssm = self.read(0xB450, 1)[0]
|
||||
if ltssm != 0x78: self.set_pcie_power(True)
|
||||
ltssm = self.read(0xB450, 1)[0]
|
||||
if ltssm != 0x78: raise RuntimeError(f"PCIe link not up (LTSSM=0x{ltssm:02X}), custom firmware not ready")
|
||||
|
||||
def set_pcie_power(self, enabled:bool, timeout:int=10000):
|
||||
checked(libusb.libusb_control_transfer,
|
||||
f"F3 PCIe power {'on' if enabled else 'off'} failed")(self.usb.handle, 0x40, 0xF3, int(enabled), 0, None, 0, timeout)
|
||||
|
||||
# === PCIe TLP via 0xF0 vendor command ===
|
||||
def set_pcie_power(self, enabled:bool, timeout:int=10000): self.usb.control_write(0xF3, value=int(enabled), timeout=timeout)
|
||||
|
||||
def _f0_out(self, fmt_type:int, byte_en:int, address:int, value:int, mode:int=0):
|
||||
struct.pack_into('<III', self._f0_out_mv, 0, address & 0xFFFFFFFF, address >> 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('<III', address & 0xFFFFFFFF, address >> 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('<I', self._f0_in_buf, 0)[0], (self._f0_in_buf[4] >> 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('<I', data)[0], (data[4] >> 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
|
||||
|
||||
Reference in New Issue
Block a user