mirror of
https://github.com/sunnypilot/sunnypilot.git
synced 2026-09-12 04:33:43 +08:00
Compare commits
203 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 5360beb1ef | |||
| d0ab77cf2e | |||
| 50d7f75bfc | |||
| 15c7f52e40 | |||
| 8dfe04a318 | |||
| 9648ac2f04 | |||
| a03a11e333 | |||
| 53b1070b09 | |||
| 68bf6d8162 | |||
| 8b472291d0 | |||
| 797c8d4293 | |||
| 7568528507 | |||
| 6665c3acf6 | |||
| 3b13eb4d0d | |||
| f7eb4c2561 | |||
| be9255358c | |||
| f48ccb5d7a | |||
| 99559d6749 | |||
| 4b6f0ffb46 | |||
| 55415382ce | |||
| 9f30901eb6 | |||
| d2045d24fb | |||
| 6488a11e4d | |||
| 7b22b65313 | |||
| b59a50101e | |||
| 3e30962bd3 | |||
| 3274043063 | |||
| c624550d2e | |||
| 6759038671 | |||
| aa636e75c8 | |||
| 5ff3c3bd8d | |||
| 1e11b52290 | |||
| 76e0204025 | |||
| cc88f2cbd6 | |||
| ee0ac199a4 | |||
| d5ed828eaa | |||
| 9fca585f2a | |||
| bb1259303e | |||
| e135051ca8 | |||
| 12bef55d8a | |||
| 2f7a45e6c8 | |||
| 936ebfc12b | |||
| aa0c9dc0eb | |||
| 7476a866e7 | |||
| 610d857e33 | |||
| d2f47407d0 | |||
| db75ec76ea | |||
| 24066465d7 | |||
| a7abbd6e25 | |||
| 878982447c | |||
| 576527a36b | |||
| ef8c35da24 | |||
| 85688b1040 | |||
| 0fb2199130 | |||
| d48d756c1d | |||
| 2ed298a0c9 | |||
| d68f038949 | |||
| 7231571e57 | |||
| b37f1419d3 | |||
| cd85a66790 | |||
| 305ea87daf | |||
| 4bbfc793e0 | |||
| d5d983676e | |||
| de8a96a398 | |||
| 0cbf45f699 | |||
| 0d68a3a2ab | |||
| 9e85a85059 | |||
| 0373c327c0 | |||
| efe9e5c200 | |||
| 8a249a45dc | |||
| bdbefe67f6 | |||
| 675bb166ad | |||
| 1b717a7e88 | |||
| 86f55a8ba9 | |||
| 629392d2f7 | |||
| bc414bdc8b | |||
| 7ca5649f2c | |||
| 641ee8fa87 | |||
| 56c276158c | |||
| c65308a8bd | |||
| 994e526460 | |||
| 1defae36b7 | |||
| 8f029fd0ef | |||
| ddb46284dc | |||
| 9effc754d9 | |||
| e49ffc2a2d | |||
| 2cacd0b3e5 | |||
| c4b8859dff | |||
| 8fb0953205 | |||
| 63d1c8835f | |||
| 17a185606d | |||
| da10131392 | |||
| 7107c2ba14 | |||
| 95b6e877ac | |||
| eb02c6570e | |||
| 1be8ae31c4 | |||
| 04dcd38856 | |||
| 22ccf0d72f | |||
| 3c969bb627 | |||
| 20f8011feb | |||
| 9cf17e74a1 | |||
| 2c4efdf557 | |||
| 4cd3d3c16c | |||
| 637f3ae9c8 | |||
| 464ee80f71 | |||
| 2743a04613 | |||
| 7f9978d001 | |||
| 4b83961c67 | |||
| c00eaf428a | |||
| 0a9993e8d4 | |||
| 0af214a985 | |||
| af43385e3a | |||
| 0ab2b8c590 | |||
| 67ab18a0de | |||
| e87dc15b30 | |||
| 192d08516c | |||
| 3cf001c59c | |||
| f2ccd021da | |||
| c9fc900f64 | |||
| 3c37c5ce5d | |||
| 7c45889e4e | |||
| 2aabb7aee8 | |||
| 3859e9962f | |||
| 810efbab72 | |||
| ec27bec326 | |||
| 250d553157 | |||
| cea54a0ca8 | |||
| 8e72d783bd | |||
| 1b0dc103dc | |||
| 6c364d292b | |||
| bcdec2ce84 | |||
| 3deaeb3759 | |||
| c669f0984a | |||
| 46dd946740 | |||
| 9da4b3653e | |||
| 4e21ae7c50 | |||
| bb91e92237 | |||
| 14b4c4f85b | |||
| 0660b542c3 | |||
| 2b893b90c9 | |||
| f5139178ed | |||
| fb43b755f2 | |||
| 07f5b967d8 | |||
| ea19c7d3bb | |||
| e461842cbb | |||
| a73c9659d5 | |||
| cb796fbc76 | |||
| 6bf75fc557 | |||
| 9a1fc28819 | |||
| 0741d05e92 | |||
| 1ad008107d | |||
| feebd9df93 | |||
| c2e5ced3e5 | |||
| 15e5d2efb9 | |||
| a3929d0b54 | |||
| 794f8f9991 | |||
| 68fa5e3f21 | |||
| 86c6cc1f48 | |||
| eb7ffbf093 | |||
| 3919095752 | |||
| 74d63be1c3 | |||
| 8894486a1a | |||
| 810599315d | |||
| 6f3ab810c8 | |||
| 230f78b8d3 | |||
| f1affec088 | |||
| 97d8ef242c | |||
| a63fff9b45 | |||
| cb3893daaa | |||
| 29f60df74b | |||
| c6c072e1f4 | |||
| d101cbb83e | |||
| 1536d59633 | |||
| dc99b865ae | |||
| e59bc027ff | |||
| cf7e5efaca | |||
| 4b44f2eb31 | |||
| 107d2ab400 | |||
| 5432d9062c | |||
| f533f6c843 | |||
| 58e9ac763c | |||
| cb50d54169 | |||
| bd5de4ed0a | |||
| 0d4073fadb | |||
| ebc70dcb52 | |||
| 4d0426999e | |||
| 286da42573 | |||
| 8a836710a9 | |||
| 5d515bcf33 | |||
| 1c7f6d5133 | |||
| 05d57c7aeb | |||
| e4b0eaf352 | |||
| a710276472 | |||
| af086db671 | |||
| 0d9eb0e25e | |||
| 0616caed6d | |||
| 095337b3c1 | |||
| 1edec2d22c | |||
| affabb9ee0 | |||
| dc27e8711c | |||
| cf7329a264 | |||
| 5ee5ecd820 | |||
| b064f730dd |
@@ -0,0 +1,11 @@
|
||||
* @sunnypilot/dev-internal
|
||||
/.github/ @devtekve @sunnyhaibin
|
||||
/release/ci/ @devtekve @sunnyhaibin
|
||||
/tinygrad_repo @devtekve @Discountchubbs
|
||||
/tinygrad/ @devtekve @Discountchubbs
|
||||
/selfdrive/controls/lib/longitudinal_planner.py @devtekve @Discountchubbs
|
||||
/selfdrive/controls/lib/longitudinal_mpc_lib/long_mpc.py @devtekve @Discountchubbs
|
||||
/selfdrive/modeld/ @devtekve @Discountchubbs
|
||||
/sunnypilot/model* @devtekve @Discountchubbs
|
||||
/sunnypilot/sunnylink/ @devtekve
|
||||
/system/athena/ @devtekve
|
||||
@@ -78,7 +78,6 @@ jobs:
|
||||
- name: Get next recompiled dir number
|
||||
id: create-recompiled-dir
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
HF_REPO: ${{ github.event.inputs.hf_repo }}
|
||||
run: |
|
||||
pip install huggingface_hub
|
||||
|
||||
@@ -30,7 +30,6 @@ jobs:
|
||||
runs-on: ubuntu-24.04
|
||||
outputs:
|
||||
model_name: ${{ steps.resolve.outputs.model_name }}
|
||||
safe_model_name: ${{ steps.resolve.outputs.safe_model_name }}
|
||||
onnx_ref: ${{ steps.resolve.outputs.onnx_ref }}
|
||||
onnx_path: ${{ steps.resolve.outputs.onnx_path }}
|
||||
hf_defaults_path: ${{ steps.resolve.outputs.hf_defaults_path }}
|
||||
@@ -65,9 +64,7 @@ jobs:
|
||||
exit 1
|
||||
fi
|
||||
|
||||
SAFE_NAME="${NAME// /-}"
|
||||
echo "model_name=${NAME}" >> $GITHUB_OUTPUT
|
||||
echo "safe_model_name=${SAFE_NAME}" >> $GITHUB_OUTPUT
|
||||
echo "onnx_ref=${ONNX_REF}" >> $GITHUB_OUTPUT
|
||||
echo "onnx_path=${ONNX_PATH}" >> $GITHUB_OUTPUT
|
||||
echo "hf_defaults_path=${HF_DEFAULTS_PATH}" >> $GITHUB_OUTPUT
|
||||
@@ -138,7 +135,7 @@ jobs:
|
||||
|
||||
- name: Prepare output
|
||||
env:
|
||||
MODEL_NAME: ${{ needs.resolve.outputs.safe_model_name }}
|
||||
MODEL_NAME: ${{ needs.resolve.outputs.model_name }}
|
||||
run: |
|
||||
source ${UV_PROJECT_ENVIRONMENT}/bin/activate
|
||||
export PYTHONPATH=${{ github.workspace }}
|
||||
@@ -161,13 +158,13 @@ jobs:
|
||||
- name: Upload small model artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: model-${{ needs.resolve.outputs.safe_model_name }}-${{ github.run_number }}
|
||||
name: model-${{ needs.resolve.outputs.model_name }}-${{ github.run_number }}
|
||||
path: ${{ github.workspace }}/small_output/
|
||||
|
||||
- name: Upload artifact name file
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: artifact-name-${{ needs.resolve.outputs.safe_model_name }}
|
||||
name: artifact-name-${{ needs.resolve.outputs.model_name }}
|
||||
path: ${{ github.workspace }}/small_output/artifact_name.txt
|
||||
|
||||
- name: Re-enable powersave
|
||||
@@ -257,7 +254,7 @@ jobs:
|
||||
|
||||
- name: Prepare output
|
||||
env:
|
||||
MODEL_NAME: ${{ needs.resolve.outputs.safe_model_name }}
|
||||
MODEL_NAME: ${{ needs.resolve.outputs.model_name }}
|
||||
run: |
|
||||
source ${UV_PROJECT_ENVIRONMENT}/bin/activate
|
||||
export PYTHONPATH=${{ github.workspace }}
|
||||
@@ -280,13 +277,13 @@ jobs:
|
||||
- name: Upload big model artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: model-${{ needs.resolve.outputs.safe_model_name }}-${{ github.run_number }}
|
||||
name: model-${{ needs.resolve.outputs.model_name }}-${{ github.run_number }}
|
||||
path: ${{ github.workspace }}/big_output/
|
||||
|
||||
- name: Upload artifact name file
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: artifact-name-${{ needs.resolve.outputs.safe_model_name }}
|
||||
name: artifact-name-${{ needs.resolve.outputs.model_name }}
|
||||
path: ${{ github.workspace }}/big_output/artifact_name.txt
|
||||
|
||||
- name: Re-enable powersave
|
||||
@@ -321,7 +318,7 @@ jobs:
|
||||
if: ${{ inputs.target == 'small' || inputs.target == 'big' }}
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: artifact-name-${{ needs.resolve.outputs.safe_model_name }}
|
||||
name: artifact-name-${{ needs.resolve.outputs.model_name }}
|
||||
path: artifact_name
|
||||
|
||||
- name: Read artifact name
|
||||
@@ -341,7 +338,7 @@ jobs:
|
||||
- name: Upload model to HF
|
||||
if: ${{ inputs.target == 'small' || inputs.target == 'big' }}
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
HF_OIDC_RESOURCE: datasets/${{ env.HF_REPO }}
|
||||
ARTIFACT_NAME: ${{ steps.artifact.outputs.artifact_name }}
|
||||
run: |
|
||||
rm -f output/artifact_name.txt
|
||||
@@ -367,7 +364,7 @@ jobs:
|
||||
- name: Generate DM metadata and upload to HF
|
||||
if: ${{ inputs.target == 'dm' }}
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
HF_OIDC_RESOURCE: datasets/${{ env.HF_REPO }}
|
||||
run: |
|
||||
export PYTHONPATH=$(pwd)
|
||||
python3 -c "
|
||||
@@ -484,29 +481,11 @@ jobs:
|
||||
print(f'Chunked {pkl} into {len(targets)} chunks')
|
||||
"
|
||||
|
||||
- name: Compile DM warp
|
||||
run: |
|
||||
source ${UV_PROJECT_ENVIRONMENT}/bin/activate
|
||||
export PYTHONPATH="${PYTHONPATH}:${{ github.workspace }}/tinygrad_repo:${{ github.workspace }}"
|
||||
|
||||
TG_FLAGS="DEV=QCOM IMAGE=1 FLOAT16=1 NOLOCALS=1 JIT_BATCH_SIZE=0 OPENPILOT_HACKS=1"
|
||||
MODEL_DIR="${{ github.workspace }}/openpilot/selfdrive/modeld"
|
||||
DM_SIZE=$(python3 -c "from openpilot.common.transformations.model import DM_INPUT_SIZE as s; print(f'{s[0]}x{s[1]}')")
|
||||
|
||||
for res in $(python3 -c "from openpilot.common.transformations.camera import _ar_ox_fisheye as a, _os_fisheye as o; print(f'{a.width}x{a.height} {o.width}x{o.height}')"); do
|
||||
WARP_PKL="${MODEL_DIR}/models/dm_warp_${res}_tinygrad.pkl"
|
||||
taskset -c 7 env ${TG_FLAGS} python3 ${MODEL_DIR}/compile_dm_warp.py \
|
||||
--camera-resolution ${res} \
|
||||
--warp-to ${DM_SIZE} \
|
||||
--output ${WARP_PKL}
|
||||
done
|
||||
|
||||
- name: Prepare DM output
|
||||
run: |
|
||||
mkdir -p dm_output
|
||||
cp ${{ github.workspace }}/${{ env.DM_PKL }}.chunk* dm_output/
|
||||
cp ${{ github.workspace }}/${{ env.DM_PKL }}.chunkmanifest dm_output/
|
||||
cp ${{ github.workspace }}/openpilot/selfdrive/modeld/models/dm_warp_* dm_output/
|
||||
|
||||
- name: Upload DM artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
|
||||
@@ -146,7 +146,7 @@ jobs:
|
||||
|
||||
- name: Validate hf_repo and JSON version
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
HF_OIDC_RESOURCE: datasets/${{ inputs.hf_repo }}
|
||||
run: |
|
||||
if [ ! -f "$JSON_FILE" ]; then
|
||||
echo "JSON file $JSON_FILE does not exist!"
|
||||
@@ -155,8 +155,13 @@ jobs:
|
||||
python3 -c "
|
||||
import sys
|
||||
from huggingface_hub import HfApi
|
||||
HfApi().repo_info(repo_id=sys.argv[1], repo_type='dataset')
|
||||
print(f'Success: Repo {sys.argv[1]} exists.')
|
||||
try:
|
||||
api = HfApi()
|
||||
api.repo_info(repo_id=sys.argv[1], repo_type='dataset')
|
||||
print(f'Success: Repo {sys.argv[1]} exists.')
|
||||
except Exception as e:
|
||||
print('HF validation failed:', e)
|
||||
sys.exit(1)
|
||||
" "${{ inputs.hf_repo }}"
|
||||
|
||||
- name: Download artifact name file
|
||||
@@ -187,7 +192,7 @@ jobs:
|
||||
|
||||
- name: Upload to Hugging Face
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
HF_OIDC_RESOURCE: datasets/${{ inputs.hf_repo }}
|
||||
ARTIFACT_NAME: ${{ steps.read-artifact-name.outputs.artifact_name }}
|
||||
run: |
|
||||
hf upload ${{ inputs.hf_repo }} \
|
||||
|
||||
@@ -46,13 +46,6 @@ runs:
|
||||
printf '%s\t%s\n' "$ENCODED_URL" "${DEST_DIR}/${CANONICAL}.chunk${CHUNK_IDX}" >> "$DOWNLOAD_LIST"
|
||||
done < <(echo "$ARTIFACT" | jq -r '.chunks[].file_name')
|
||||
echo "$NUM_CHUNKS" > "${DEST_DIR}/${CANONICAL}.chunkmanifest"
|
||||
|
||||
if [ "$CANONICAL" = "dmonitoring_model_tinygrad.pkl" ]; then
|
||||
for warp in dm_warp_1928x1208_tinygrad.pkl dm_warp_1344x760_tinygrad.pkl; do
|
||||
ENCODED_URL=$(python3 -c "import urllib.parse; print(urllib.parse.quote('${BASE_URL}/${warp}', safe=':/'))")
|
||||
printf '%s\t%s\n' "$ENCODED_URL" "${DEST_DIR}/${warp}" >> "$DOWNLOAD_LIST"
|
||||
done
|
||||
fi
|
||||
}
|
||||
|
||||
echo "$MODELS_JSON" | jq -c '.[]' | while IFS= read -r model; do
|
||||
|
||||
@@ -188,7 +188,7 @@ jobs:
|
||||
if [ "${{ inputs.target_hardware }}" == "chestnut" ]; then
|
||||
echo "CHESTNUT build"
|
||||
export CHESTNUT=1
|
||||
TG_FLAGS="DEBUG=1 DEV=USB+AMD:LLVM FLOAT16=1 JIT_BATCH_SIZE=0 GMMU=0 TC_OPT=2 TC_OCCUPANCY_OPT=1"
|
||||
TG_FLAGS="DEBUG=1 DEV=USB+AMD:LLVM WARP_DEV=QCOM FLOAT16=1 JIT_BATCH_SIZE=0 GMMU=0 TC_OPT=2"
|
||||
OUTPUT_PKL="${{ env.MODELS_DIR }}/big_driving_tinygrad.pkl"
|
||||
else
|
||||
echo "QCOM build"
|
||||
|
||||
@@ -216,9 +216,6 @@ jobs:
|
||||
needs: [ prepare_strategy ]
|
||||
runs-on: ubuntu-24.04
|
||||
if: ${{ needs.prepare_strategy.outputs.include_big_model == 'true' }}
|
||||
concurrency:
|
||||
group: prepare-chestnut
|
||||
cancel-in-progress: false
|
||||
outputs:
|
||||
onnx_sha256: ${{ steps.resolve.outputs.onnx_sha256 }}
|
||||
env:
|
||||
@@ -231,10 +228,8 @@ jobs:
|
||||
run: |
|
||||
REF="${{ github.head_ref || github.ref_name }}"
|
||||
|
||||
BLOB_SHA=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx?ref=${REF}" --jq '.sha')
|
||||
ONNX_HASH=$(gh api "repos/${GH_REPO}/git/blobs/${BLOB_SHA}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
ONNX_HASH=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx?ref=${REF}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
echo "ONNX hash: $ONNX_HASH"
|
||||
[ -n "$ONNX_HASH" ] || { echo "::error::Failed to extract ONNX hash"; exit 1; }
|
||||
echo "onnx_sha256=$ONNX_HASH" >> $GITHUB_OUTPUT
|
||||
|
||||
TINYGRAD_REF=$(gh api "repos/${GH_REPO}/contents/tinygrad_repo?ref=${REF}" --jq '.sha')
|
||||
@@ -243,7 +238,7 @@ jobs:
|
||||
JSON_URL="https://huggingface.co/datasets/${HF_REPO}/resolve/main/${HF_DEFAULTS_PATH}/default_models.json"
|
||||
|
||||
check_defaults() {
|
||||
DEFAULTS=$(curl -fsSL "${JSON_URL}?t=$(date +%s)" 2>/dev/null) || return 1
|
||||
DEFAULTS=$(curl -fsSL "$JSON_URL" 2>/dev/null) || return 1
|
||||
TINYGRAD_MATCH=$(echo "$DEFAULTS" | jq -r --arg ref "$TINYGRAD_REF" '.tinygrad_ref == $ref' 2>/dev/null)
|
||||
[ "$TINYGRAD_MATCH" = "true" ] || return 1
|
||||
BUNDLE=$(echo "$DEFAULTS" | jq --arg hash "$ONNX_HASH" '.bundles[] | select(.onnx_sha256 == $hash)' 2>/dev/null)
|
||||
@@ -257,35 +252,18 @@ jobs:
|
||||
|
||||
echo "No matching model on HF — dispatching build"
|
||||
gh workflow run build-default-models.yaml --ref "$REF" -f target=big
|
||||
sleep 10
|
||||
|
||||
BUILD_RUN_ID=$(gh run list --workflow build-default-models.yaml --branch "$REF" --limit 1 --json databaseId --jq '.[0].databaseId')
|
||||
echo "Dispatched build run: $BUILD_RUN_ID"
|
||||
|
||||
echo "Waiting for build run to complete..."
|
||||
echo "Polling HF for big model availability..."
|
||||
for i in $(seq 1 90); do
|
||||
sleep 30
|
||||
STATUS=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.status')
|
||||
CONCLUSION=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.conclusion')
|
||||
echo "Poll $i/90: status=$STATUS conclusion=$CONCLUSION"
|
||||
if [ "$STATUS" = "completed" ]; then
|
||||
if [ "$CONCLUSION" = "success" ]; then
|
||||
echo "Build run succeeded, verifying HF..."
|
||||
sleep 10
|
||||
if check_defaults; then
|
||||
echo "Big model verified on HF"
|
||||
exit 0
|
||||
fi
|
||||
echo "::error::Build succeeded but model not found on HF"
|
||||
exit 1
|
||||
else
|
||||
echo "::error::Build run failed with conclusion=$CONCLUSION"
|
||||
exit 1
|
||||
fi
|
||||
if check_defaults; then
|
||||
echo "Big model available on HF after $((i * 30))s"
|
||||
exit 0
|
||||
fi
|
||||
echo "Poll $i/90: not yet available"
|
||||
done
|
||||
|
||||
echo "::error::Build run did not complete within 45 minutes"
|
||||
echo "::error::Big model not available on HF after 45 minutes"
|
||||
exit 1
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
@@ -299,9 +277,6 @@ jobs:
|
||||
prepare_small_model:
|
||||
needs: [ prepare_strategy ]
|
||||
runs-on: ubuntu-24.04
|
||||
concurrency:
|
||||
group: prepare-small-model
|
||||
cancel-in-progress: false
|
||||
outputs:
|
||||
driving_onnx_sha256: ${{ steps.resolve.outputs.driving_onnx_sha256 }}
|
||||
env:
|
||||
@@ -314,10 +289,8 @@ jobs:
|
||||
run: |
|
||||
REF="${{ github.head_ref || github.ref_name }}"
|
||||
|
||||
BLOB_SHA=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/driving_supercombo.onnx?ref=${REF}" --jq '.sha')
|
||||
DRIVING_HASH=$(gh api "repos/${GH_REPO}/git/blobs/${BLOB_SHA}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
DRIVING_HASH=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/driving_supercombo.onnx?ref=${REF}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
echo "Driving ONNX hash: $DRIVING_HASH"
|
||||
[ -n "$DRIVING_HASH" ] || { echo "::error::Failed to extract driving ONNX hash"; exit 1; }
|
||||
echo "driving_onnx_sha256=$DRIVING_HASH" >> $GITHUB_OUTPUT
|
||||
|
||||
TINYGRAD_REF=$(gh api "repos/${GH_REPO}/contents/tinygrad_repo?ref=${REF}" --jq '.sha')
|
||||
@@ -326,7 +299,7 @@ jobs:
|
||||
JSON_URL="https://huggingface.co/datasets/${HF_REPO}/resolve/main/${HF_DEFAULTS_PATH}/default_models.json"
|
||||
|
||||
check_defaults() {
|
||||
DEFAULTS=$(curl -fsSL "${JSON_URL}?t=$(date +%s)" 2>/dev/null) || return 1
|
||||
DEFAULTS=$(curl -fsSL "$JSON_URL" 2>/dev/null) || return 1
|
||||
TINYGRAD_MATCH=$(echo "$DEFAULTS" | jq -r --arg ref "$TINYGRAD_REF" '.tinygrad_ref == $ref' 2>/dev/null)
|
||||
[ "$TINYGRAD_MATCH" = "true" ] || return 1
|
||||
DRIVING=$(echo "$DEFAULTS" | jq --arg hash "$DRIVING_HASH" '.bundles[] | select(.onnx_sha256 == $hash)' 2>/dev/null)
|
||||
@@ -340,35 +313,18 @@ jobs:
|
||||
|
||||
echo "No matching model on HF — dispatching build"
|
||||
gh workflow run build-default-models.yaml --ref "$REF" -f target=small
|
||||
sleep 10
|
||||
|
||||
BUILD_RUN_ID=$(gh run list --workflow build-default-models.yaml --branch "$REF" --limit 1 --json databaseId --jq '.[0].databaseId')
|
||||
echo "Dispatched build run: $BUILD_RUN_ID"
|
||||
|
||||
echo "Waiting for build run to complete..."
|
||||
echo "Polling HF for model availability..."
|
||||
for i in $(seq 1 60); do
|
||||
sleep 30
|
||||
STATUS=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.status')
|
||||
CONCLUSION=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.conclusion')
|
||||
echo "Poll $i/60: status=$STATUS conclusion=$CONCLUSION"
|
||||
if [ "$STATUS" = "completed" ]; then
|
||||
if [ "$CONCLUSION" = "success" ]; then
|
||||
echo "Build run succeeded, verifying HF..."
|
||||
sleep 10
|
||||
if check_defaults; then
|
||||
echo "Small model verified on HF"
|
||||
exit 0
|
||||
fi
|
||||
echo "::error::Build succeeded but model not found on HF"
|
||||
exit 1
|
||||
else
|
||||
echo "::error::Build run failed with conclusion=$CONCLUSION"
|
||||
exit 1
|
||||
fi
|
||||
if check_defaults; then
|
||||
echo "Model available on HF after $((i * 30))s"
|
||||
exit 0
|
||||
fi
|
||||
echo "Poll $i/60: not yet available"
|
||||
done
|
||||
|
||||
echo "::error::Small model build did not complete within 30 minutes"
|
||||
echo "::error::Small driving model not available on HF after 30 minutes"
|
||||
exit 1
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
@@ -382,9 +338,6 @@ jobs:
|
||||
prepare_dm_model:
|
||||
needs: [ prepare_strategy ]
|
||||
runs-on: ubuntu-24.04
|
||||
concurrency:
|
||||
group: prepare-dm-model
|
||||
cancel-in-progress: false
|
||||
outputs:
|
||||
dm_onnx_sha256: ${{ steps.resolve.outputs.dm_onnx_sha256 }}
|
||||
env:
|
||||
@@ -397,10 +350,8 @@ jobs:
|
||||
run: |
|
||||
REF="${{ github.head_ref || github.ref_name }}"
|
||||
|
||||
BLOB_SHA=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/dmonitoring_model.onnx?ref=${REF}" --jq '.sha')
|
||||
DM_HASH=$(gh api "repos/${GH_REPO}/git/blobs/${BLOB_SHA}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
DM_HASH=$(gh api "repos/${GH_REPO}/contents/openpilot/selfdrive/modeld/models/dmonitoring_model.onnx?ref=${REF}" --jq '.content' | base64 -d | grep '^oid sha256:' | cut -d: -f2)
|
||||
echo "DM ONNX hash: $DM_HASH"
|
||||
[ -n "$DM_HASH" ] || { echo "::error::Failed to extract DM ONNX hash"; exit 1; }
|
||||
echo "dm_onnx_sha256=$DM_HASH" >> $GITHUB_OUTPUT
|
||||
|
||||
TINYGRAD_REF=$(gh api "repos/${GH_REPO}/contents/tinygrad_repo?ref=${REF}" --jq '.sha')
|
||||
@@ -409,7 +360,7 @@ jobs:
|
||||
JSON_URL="https://huggingface.co/datasets/${HF_REPO}/resolve/main/${HF_DEFAULTS_PATH}/default_models.json"
|
||||
|
||||
check_defaults() {
|
||||
DEFAULTS=$(curl -fsSL "${JSON_URL}?t=$(date +%s)" 2>/dev/null) || return 1
|
||||
DEFAULTS=$(curl -fsSL "$JSON_URL" 2>/dev/null) || return 1
|
||||
TINYGRAD_MATCH=$(echo "$DEFAULTS" | jq -r --arg ref "$TINYGRAD_REF" '.tinygrad_ref == $ref' 2>/dev/null)
|
||||
[ "$TINYGRAD_MATCH" = "true" ] || return 1
|
||||
DM=$(echo "$DEFAULTS" | jq --arg hash "$DM_HASH" '.bundles[] | select(.onnx_sha256 == $hash)' 2>/dev/null)
|
||||
@@ -423,35 +374,18 @@ jobs:
|
||||
|
||||
echo "No matching DM model on HF — dispatching build"
|
||||
gh workflow run build-default-models.yaml --ref "$REF" -f target=dm
|
||||
sleep 10
|
||||
|
||||
BUILD_RUN_ID=$(gh run list --workflow build-default-models.yaml --branch "$REF" --limit 1 --json databaseId --jq '.[0].databaseId')
|
||||
echo "Dispatched build run: $BUILD_RUN_ID"
|
||||
|
||||
echo "Waiting for build run to complete..."
|
||||
echo "Polling HF for DM model availability..."
|
||||
for i in $(seq 1 60); do
|
||||
sleep 30
|
||||
STATUS=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.status')
|
||||
CONCLUSION=$(gh api "repos/${GH_REPO}/actions/runs/${BUILD_RUN_ID}" --jq '.conclusion')
|
||||
echo "Poll $i/60: status=$STATUS conclusion=$CONCLUSION"
|
||||
if [ "$STATUS" = "completed" ]; then
|
||||
if [ "$CONCLUSION" = "success" ]; then
|
||||
echo "Build run succeeded, verifying HF..."
|
||||
sleep 10
|
||||
if check_defaults; then
|
||||
echo "DM model verified on HF"
|
||||
exit 0
|
||||
fi
|
||||
echo "::error::Build succeeded but DM model not found on HF"
|
||||
exit 1
|
||||
else
|
||||
echo "::error::Build run failed with conclusion=$CONCLUSION"
|
||||
exit 1
|
||||
fi
|
||||
if check_defaults; then
|
||||
echo "DM model available on HF after $((i * 30))s"
|
||||
exit 0
|
||||
fi
|
||||
echo "Poll $i/60: not yet available"
|
||||
done
|
||||
|
||||
echo "::error::DM model build did not complete within 30 minutes"
|
||||
echo "::error::DM model not available on HF after 30 minutes"
|
||||
exit 1
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
+1
-1
Submodule opendbc_repo updated: f95f996f59...0819b0e8e0
@@ -227,11 +227,6 @@ inline static std::unordered_map<std::string, ParamKeyAttributes> keys = {
|
||||
{"SunnylinkEnabled", {PERSISTENT, BOOL, "1"}},
|
||||
{"SunnylinkTempFault", {CLEAR_ON_MANAGER_START | CLEAR_ON_OFFROAD_TRANSITION, BOOL, "0"}},
|
||||
|
||||
{"SunnylinkLocalApps", {PERSISTENT, JSON}},
|
||||
{"SunnylinkLocalPairingCode", {CLEAR_ON_MANAGER_START, JSON}},
|
||||
{"SunnylinkLocalDiscoveredApp", {CLEAR_ON_MANAGER_START, JSON}},
|
||||
{"SunnylinkLocalPairingRequest", {CLEAR_ON_MANAGER_START, BOOL}},
|
||||
|
||||
// Backup Manager params
|
||||
{"BackupManager_CreateBackup", {PERSISTENT, BOOL}},
|
||||
{"BackupManager_RestoreVersion", {PERSISTENT, STRING}},
|
||||
|
||||
@@ -5,42 +5,22 @@ This file is part of sunnypilot and is licensed under the MIT License.
|
||||
See the LICENSE.md file in the root directory for more details.
|
||||
"""
|
||||
import pyray as rl
|
||||
from functools import partial
|
||||
from openpilot.cereal import custom
|
||||
from openpilot.common.version import sunnylink_consent_version
|
||||
from openpilot.selfdrive.ui.sunnypilot.layouts.onboarding import SunnylinkConsentPage
|
||||
from openpilot.selfdrive.ui.ui_state import ui_state
|
||||
from openpilot.sunnypilot.sunnylink.api import UNREGISTERED_SUNNYLINK_DONGLE_ID
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_discovery import latest_discovered_app
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_pairing import (
|
||||
LocalApp,
|
||||
arm_pairing,
|
||||
clear_pairing_request,
|
||||
get_local_apps,
|
||||
local_app_display_name,
|
||||
pairing_requested,
|
||||
read_pairing_code,
|
||||
remove_local_app,
|
||||
)
|
||||
from openpilot.system.ui.lib.application import gui_app, FontWeight, TextAlignment, TextAlignmentVertical
|
||||
from openpilot.system.ui.lib.multilang import tr
|
||||
from openpilot.system.ui.lib.text_measure import measure_text_cached
|
||||
from openpilot.system.ui.lib.wrap_text import wrap_text
|
||||
from openpilot.system.ui.sunnypilot.widgets.list_view import ListItemSP, button_item_sp, toggle_item_sp
|
||||
from openpilot.system.ui.sunnypilot.widgets.list_view import button_item_sp
|
||||
from openpilot.system.ui.sunnypilot.widgets.list_view import toggle_item_sp
|
||||
from openpilot.system.ui.sunnypilot.widgets.sunnylink_pairing_dialog import SunnylinkPairingDialog
|
||||
from openpilot.system.ui.widgets import Widget, DialogResult
|
||||
from openpilot.system.ui.widgets.button import ButtonStyle, Button, IconButton
|
||||
from openpilot.system.ui.widgets.button import ButtonStyle, Button
|
||||
from openpilot.system.ui.widgets.confirm_dialog import alert_dialog, ConfirmDialog
|
||||
from openpilot.system.ui.widgets.label import UnifiedLabel
|
||||
from openpilot.system.ui.widgets.list_view import dual_button_item
|
||||
from openpilot.system.ui.widgets.network import NavButton
|
||||
from openpilot.system.ui.widgets.scroller_tici import Scroller, LineSeparator
|
||||
|
||||
MAX_LOCAL_APPS = 4
|
||||
|
||||
# Read-only value colors used by the local-mode rows.
|
||||
_LOCAL_DISCOVERED_COLOR = rl.Color(170, 170, 170, 255) # grey: no app in sight
|
||||
_LOCAL_ACTIVE_COLOR = rl.Color(0, 255, 0, 255) # green: discovered / pairing code
|
||||
from openpilot.common.version import sunnylink_consent_version
|
||||
|
||||
|
||||
class SunnylinkHeader(Widget):
|
||||
@@ -212,15 +192,6 @@ class SunnylinkLayout(Widget):
|
||||
self._backup_btn.set_button_style(ButtonStyle.NORMAL)
|
||||
self._restore_btn.set_button_style(ButtonStyle.PRIMARY)
|
||||
|
||||
self._mobile_app_btn = button_item_sp(
|
||||
title=tr("Sunnylink Local Connections"),
|
||||
button_text=tr("CONFIGURE"),
|
||||
description=tr("Manage the mobile app(s) connected over Wi-Fi: pair a new app ") +
|
||||
tr("or unpair existing ones."),
|
||||
callback=self._open_local_apps,
|
||||
)
|
||||
self._mobile_app_btn.set_visible(lambda: self._sunnylink_enabled)
|
||||
|
||||
items = [
|
||||
SunnylinkHeader(),
|
||||
LineSeparator(),
|
||||
@@ -231,11 +202,9 @@ class SunnylinkLayout(Widget):
|
||||
LineSeparator(),
|
||||
self._pair_btn,
|
||||
LineSeparator(),
|
||||
self._mobile_app_btn,
|
||||
LineSeparator(),
|
||||
self._sunnylink_uploader_toggle,
|
||||
LineSeparator(),
|
||||
self._sunnylink_backup_restore_buttons,
|
||||
self._sunnylink_backup_restore_buttons
|
||||
]
|
||||
return items
|
||||
|
||||
@@ -348,8 +317,6 @@ class SunnylinkLayout(Widget):
|
||||
gui_app.push_widget(sl_terms_dlg)
|
||||
else:
|
||||
ui_state.params.put_bool("SunnylinkEnabled", state)
|
||||
if not state:
|
||||
clear_pairing_request()
|
||||
self._update_description(state)
|
||||
|
||||
def _update_description(self, state: bool):
|
||||
@@ -385,9 +352,6 @@ class SunnylinkLayout(Widget):
|
||||
self._pair_btn.action_item.set_text(pair_btn_text)
|
||||
self._pair_btn.action_item.set_enabled(self._sunnylink_enabled)
|
||||
|
||||
def _open_local_apps(self):
|
||||
gui_app.push_widget(SunnylinkLocalAppLayout())
|
||||
|
||||
def _render(self, rect):
|
||||
self._scroller.render(rect)
|
||||
|
||||
@@ -400,157 +364,3 @@ class SunnylinkLayout(Widget):
|
||||
def hide_event(self):
|
||||
super().hide_event()
|
||||
ui_state.sunnylink_state.set_settings_open(False)
|
||||
|
||||
|
||||
class SunnylinkLocalAppLayout(Widget):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._local_apps_cache: list[LocalApp] = []
|
||||
|
||||
self._back_button = NavButton(tr("Back"))
|
||||
self._back_button.set_click_callback(gui_app.pop_widget)
|
||||
|
||||
self._pair_app_btn = button_item_sp(
|
||||
title=tr("Pair App"),
|
||||
button_text=tr("PAIR"),
|
||||
description=tr("Open a 5-minute pairing window and show the code to ") +
|
||||
tr("type into the app. Closing the dialog cancels pairing."),
|
||||
callback=self._show_pairing_code_dialog,
|
||||
)
|
||||
|
||||
self._local_app_rows: list[ListItemSP] = []
|
||||
self._local_app_seps: list[LineSeparator] = []
|
||||
for i in range(MAX_LOCAL_APPS):
|
||||
row = button_item_sp(
|
||||
title=lambda i=i: self._local_app_title(i),
|
||||
button_text=tr("UNPAIR"),
|
||||
description=lambda i=i: self._local_app_endpoint(i),
|
||||
callback=partial(self._unpair_local_app, i),
|
||||
)
|
||||
sep = LineSeparator()
|
||||
row.set_visible(lambda i=i: self._local_row_visible(i))
|
||||
sep.set_visible(lambda i=i: self._local_row_visible(i))
|
||||
self._local_app_rows.append(row)
|
||||
self._local_app_seps.append(sep)
|
||||
|
||||
items = [self._pair_app_btn, LineSeparator()]
|
||||
for row, sep in zip(self._local_app_rows, self._local_app_seps, strict=True):
|
||||
items.extend((row, sep))
|
||||
self._scroller = Scroller(items, line_separator=False, spacing=0)
|
||||
|
||||
def _local_row_visible(self, i: int) -> bool:
|
||||
return i < len(self._local_apps_cache)
|
||||
|
||||
def _local_app_title(self, i: int) -> str:
|
||||
if i >= len(self._local_apps_cache):
|
||||
return ""
|
||||
return local_app_display_name(self._local_apps_cache[i])
|
||||
|
||||
def _local_app_endpoint(self, i: int) -> str:
|
||||
if i >= len(self._local_apps_cache):
|
||||
return ""
|
||||
return self._local_apps_cache[i].endpoint
|
||||
|
||||
def _show_pairing_code_dialog(self):
|
||||
gui_app.push_widget(SunnylinkLocalPairingDialog())
|
||||
|
||||
def _unpair_local_app(self, index: int):
|
||||
apps = self._local_apps_cache
|
||||
if index >= len(apps):
|
||||
return
|
||||
app = apps[index]
|
||||
name = local_app_display_name(app)
|
||||
|
||||
def on_confirm(_dialog_result: int):
|
||||
remove_local_app(app.app_id)
|
||||
|
||||
dialog = ConfirmDialog(
|
||||
text=tr("Unpair") + f" {name}? " + tr("You will need the pairing code again to reconnect it."),
|
||||
confirm_text=tr("Unpair"),
|
||||
callback=on_confirm,
|
||||
)
|
||||
gui_app.push_widget(dialog)
|
||||
|
||||
def _update_state(self):
|
||||
super()._update_state()
|
||||
self._local_apps_cache = get_local_apps()
|
||||
|
||||
def _render(self, rect):
|
||||
self._back_button.set_position(self._rect.x, self._rect.y + 20)
|
||||
self._back_button.render()
|
||||
content_rect = rl.Rectangle(rect.x, rect.y + self._back_button.rect.height + 40,
|
||||
rect.width, rect.height - self._back_button.rect.height - 40)
|
||||
self._scroller.render(content_rect)
|
||||
|
||||
def show_event(self):
|
||||
super().show_event()
|
||||
self._scroller.show_event()
|
||||
|
||||
def hide_event(self):
|
||||
super().hide_event()
|
||||
self._scroller.hide_event()
|
||||
|
||||
|
||||
class SunnylinkLocalPairingDialog(Widget):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._apps_before = len(get_local_apps())
|
||||
arm_pairing()
|
||||
self._close_btn = IconButton(gui_app.texture("icons/close.png", 80, 80))
|
||||
self._close_btn.set_click_callback(self._cancel)
|
||||
|
||||
def _cancel(self):
|
||||
clear_pairing_request()
|
||||
gui_app.pop_widget()
|
||||
|
||||
def _update_state(self):
|
||||
if len(get_local_apps()) > self._apps_before:
|
||||
gui_app.pop_widget() # paired — window already cleared
|
||||
elif not pairing_requested():
|
||||
gui_app.pop_widget() # window expired
|
||||
|
||||
def _render(self, rect) -> int:
|
||||
rl.clear_background(rl.Color(224, 224, 224, 255))
|
||||
|
||||
margin = 70
|
||||
content_rect = rl.Rectangle(rect.x + margin, rect.y + margin,
|
||||
rect.width - 2 * margin, rect.height - 2 * margin)
|
||||
y = content_rect.y
|
||||
|
||||
close_size = 80
|
||||
pad = 20
|
||||
close_rect = rl.Rectangle(content_rect.x - pad, y - pad, close_size + pad * 2, close_size + pad * 2)
|
||||
self._close_btn.render(close_rect)
|
||||
y += close_size + 40
|
||||
|
||||
title_font = gui_app.font(FontWeight.NORMAL)
|
||||
title_wrapped = wrap_text(title_font, tr("Pair with mobile app"), 75, int(content_rect.width))
|
||||
rl.draw_text_ex(title_font, "\n".join(title_wrapped), rl.Vector2(content_rect.x, y), 75, 0.0, rl.BLACK)
|
||||
y += len(title_wrapped) * 75 + 40
|
||||
|
||||
code = read_pairing_code() or "—"
|
||||
code_font = gui_app.font(FontWeight.BOLD)
|
||||
code_size = measure_text_cached(code_font, code, 110)
|
||||
rl.draw_text_ex(code_font, code, rl.Vector2(content_rect.x + (content_rect.width - code_size.x) / 2, y),
|
||||
110, 0.0, rl.BLACK)
|
||||
y += 170
|
||||
|
||||
hint_font = gui_app.font(FontWeight.NORMAL)
|
||||
hint_wrapped = wrap_text(hint_font, tr("Enter this code in the sunnylink app on your phone."), 45,
|
||||
int(content_rect.width))
|
||||
rl.draw_text_ex(hint_font, "\n".join(hint_wrapped), rl.Vector2(content_rect.x, y), 45, 0.0, rl.BLACK)
|
||||
y += len(hint_wrapped) * 45 + 30
|
||||
|
||||
discovered = latest_discovered_app()
|
||||
if discovered is not None:
|
||||
endpoint, age = discovered
|
||||
status = endpoint if age < 2 else f"{endpoint} ({age}s)"
|
||||
color = _LOCAL_ACTIVE_COLOR
|
||||
else:
|
||||
status = tr("Waiting for the app…")
|
||||
color = _LOCAL_DISCOVERED_COLOR
|
||||
status_font = gui_app.font(FontWeight.NORMAL)
|
||||
rl.draw_text_ex(status_font, status, rl.Vector2(content_rect.x, y), 40, 0.0, color)
|
||||
return -1
|
||||
|
||||
@@ -5,34 +5,21 @@ This file is part of sunnypilot and is licensed under the MIT License.
|
||||
See the LICENSE.md file in the root directory for more details.
|
||||
"""
|
||||
import pyray as rl
|
||||
from functools import partial
|
||||
|
||||
|
||||
from openpilot.cereal import custom
|
||||
from openpilot.common.version import sunnylink_consent_version, sunnylink_consent_declined
|
||||
from openpilot.selfdrive.ui.mici.widgets.button import BigButton, BigToggle
|
||||
from openpilot.selfdrive.ui.mici.widgets.dialog import BigDialog, BigConfirmationDialog, BigDialogBase
|
||||
from openpilot.selfdrive.ui.mici.widgets.dialog import BigDialog, BigConfirmationDialog
|
||||
from openpilot.selfdrive.ui.sunnypilot.mici.layouts.onboarding import SunnylinkConsentPage
|
||||
from openpilot.selfdrive.ui.sunnypilot.mici.widgets.sunnylink_pairing_dialog import SunnylinkPairingDialog
|
||||
from openpilot.selfdrive.ui.ui_state import ui_state
|
||||
from openpilot.sunnypilot.sunnylink.api import UNREGISTERED_SUNNYLINK_DONGLE_ID
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_discovery import latest_discovered_app
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_pairing import (
|
||||
LocalApp,
|
||||
arm_pairing,
|
||||
clear_pairing_request,
|
||||
get_local_apps,
|
||||
local_app_display_name,
|
||||
pairing_requested,
|
||||
read_pairing_code,
|
||||
remove_local_app,
|
||||
)
|
||||
from openpilot.system.ui.lib.application import gui_app, MousePos, FontWeight
|
||||
from openpilot.system.ui.lib.multilang import tr
|
||||
from openpilot.system.ui.widgets import Widget
|
||||
from openpilot.system.ui.widgets.label import UnifiedLabel
|
||||
from openpilot.system.ui.widgets.scroller import NavScroller
|
||||
|
||||
MAX_LOCAL_APPS = 4
|
||||
from openpilot.common.version import sunnylink_consent_version, sunnylink_consent_declined
|
||||
|
||||
class SunnylinkInfo(Widget):
|
||||
def __init__(self):
|
||||
@@ -86,15 +73,11 @@ class SunnylinkLayoutMici(NavScroller):
|
||||
self._sunnylink_uploader_toggle = BigToggle(text=tr("sunnylink uploader"), initial_state=False,
|
||||
toggle_callback=self._sunnylink_uploader_callback)
|
||||
|
||||
self._mobile_app_btn = BigButton(tr("sunnylink local"), "")
|
||||
self._mobile_app_btn.set_click_callback(lambda: gui_app.push_widget(LocalAppsPanelMici()))
|
||||
|
||||
self._scroller.add_widgets([
|
||||
self._sunnylink_info,
|
||||
self._sunnylink_toggle,
|
||||
self._sunnylink_sponsor_button,
|
||||
self._sunnylink_pair_button,
|
||||
self._mobile_app_btn,
|
||||
self._backup_btn,
|
||||
self._restore_btn,
|
||||
self._sunnylink_uploader_toggle
|
||||
@@ -127,7 +110,6 @@ class SunnylinkLayoutMici(NavScroller):
|
||||
self._sunnylink_pair_button.set_text(tr("paired"))
|
||||
else:
|
||||
self._sunnylink_pair_button.set_text(tr("pair"))
|
||||
self._mobile_app_btn.set_visible(self._sunnylink_enabled)
|
||||
|
||||
def show_event(self):
|
||||
super().show_event()
|
||||
@@ -158,8 +140,6 @@ class SunnylinkLayoutMici(NavScroller):
|
||||
gui_app.push_widget(sl_terms_dlg)
|
||||
else:
|
||||
ui_state.params.put_bool("SunnylinkEnabled", state)
|
||||
if not state:
|
||||
clear_pairing_request()
|
||||
|
||||
ui_state.update_params()
|
||||
|
||||
@@ -272,108 +252,3 @@ class SunnylinkPairBigButton(BigButton):
|
||||
dlg = SunnylinkPairingDialog(sponsor_pairing=False)
|
||||
if dlg:
|
||||
gui_app.push_widget(dlg)
|
||||
|
||||
|
||||
class LocalAppsPanelMici(NavScroller):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._local_apps_cache: list[LocalApp] = []
|
||||
|
||||
self._pair_app_btn = BigButton(tr("pair app"), "")
|
||||
self._pair_app_btn.set_click_callback(lambda: gui_app.push_widget(LocalPairingCodeDialogMici()))
|
||||
|
||||
self._local_app_btns: list[BigButton] = []
|
||||
for i in range(MAX_LOCAL_APPS):
|
||||
btn = BigButton("", "")
|
||||
btn.set_click_callback(partial(self._confirm_unpair_local_app, i))
|
||||
self._local_app_btns.append(btn)
|
||||
|
||||
self._scroller.add_widgets([self._pair_app_btn, *self._local_app_btns])
|
||||
|
||||
def _update_state(self):
|
||||
super()._update_state()
|
||||
self._local_apps_cache = get_local_apps()
|
||||
for i, btn in enumerate(self._local_app_btns):
|
||||
btn.set_visible(i < len(self._local_apps_cache))
|
||||
if i < len(self._local_apps_cache):
|
||||
app = self._local_apps_cache[i]
|
||||
btn.set_text(local_app_display_name(app))
|
||||
btn.set_value(app.endpoint)
|
||||
|
||||
def _confirm_unpair_local_app(self, index: int):
|
||||
apps = self._local_apps_cache
|
||||
if index >= len(apps):
|
||||
return
|
||||
app = apps[index]
|
||||
|
||||
def unpair():
|
||||
remove_local_app(app.app_id)
|
||||
|
||||
icon = gui_app.texture("icons_mici/settings/device/update.png", 64, 64)
|
||||
dlg = BigConfirmationDialog(
|
||||
tr("slide to unpair"),
|
||||
icon,
|
||||
confirm_callback=unpair,
|
||||
red=True,
|
||||
)
|
||||
gui_app.push_widget(dlg)
|
||||
|
||||
|
||||
class LocalPairingCodeDialogMici(BigDialogBase):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self._apps_before = len(get_local_apps())
|
||||
arm_pairing()
|
||||
self.set_back_callback(clear_pairing_request)
|
||||
|
||||
header_color = rl.Color(255, 255, 255, int(255 * 0.9))
|
||||
subheader_color = rl.Color(255, 255, 255, int(255 * 0.9 * 0.65))
|
||||
self._title = UnifiedLabel(tr("pair with mobile app"), font_size=48, font_weight=FontWeight.BOLD,
|
||||
text_color=header_color, line_height=0.8)
|
||||
self._code_label = UnifiedLabel("", font_size=110, font_weight=FontWeight.DISPLAY,
|
||||
text_color=rl.Color(0, 255, 0, 255))
|
||||
self._hint = UnifiedLabel(tr("enter this code in the sunnylink app"), font_size=32,
|
||||
text_color=subheader_color, line_height=0.9)
|
||||
self._status = UnifiedLabel("", font_size=28,
|
||||
text_color=rl.Color(255, 255, 255, int(255 * 0.45)), line_height=0.9)
|
||||
|
||||
def _update_state(self):
|
||||
super()._update_state()
|
||||
if self.is_dismissing:
|
||||
return
|
||||
if len(get_local_apps()) > self._apps_before:
|
||||
self.dismiss() # paired — window already cleared
|
||||
elif not pairing_requested():
|
||||
self.dismiss() # window expired
|
||||
|
||||
def _render(self, _):
|
||||
self._code_label.set_text(read_pairing_code() or "—")
|
||||
|
||||
discovered = latest_discovered_app()
|
||||
if discovered is not None:
|
||||
endpoint, age = discovered
|
||||
self._status.set_text(endpoint if age < 2 else f"{endpoint} ({age}s)")
|
||||
self._status.set_text_color(rl.Color(0, 255, 0, 255))
|
||||
else:
|
||||
self._status.set_text(tr("waiting for the app…"))
|
||||
self._status.set_text_color(rl.Color(255, 255, 255, int(255 * 0.45)))
|
||||
|
||||
x = self._rect.x + 20
|
||||
width = int(self._rect.width - 40)
|
||||
self._title.set_max_width(width)
|
||||
self._title.set_position(x, self._rect.y + 40)
|
||||
self._title.render()
|
||||
|
||||
self._code_label.set_max_width(width)
|
||||
self._code_label.set_position(x, self._rect.y + 130)
|
||||
self._code_label.render()
|
||||
|
||||
self._hint.set_max_width(width)
|
||||
self._hint.set_position(x, self._rect.y + 290)
|
||||
self._hint.render()
|
||||
|
||||
self._status.set_max_width(width)
|
||||
self._status.set_position(x, self._rect.y + 360)
|
||||
self._status.render()
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from pathlib import Path
|
||||
|
||||
MODEL_PATH = Path(__file__).parent / 'models/supercombo.onnx'
|
||||
MODEL_PKL_PATH = Path(__file__).parent / 'models/supercombo_tinygrad.pkl'
|
||||
METADATA_PATH = Path(__file__).parent / 'models/supercombo_metadata.pkl'
|
||||
|
||||
@@ -32,7 +32,7 @@ def _patch_tinygrad_fetch_fw():
|
||||
helpers.fetch_fw = fetch_fw
|
||||
_patch_tinygrad_fetch_fw()
|
||||
|
||||
import openpilot.selfdrive.modeld.compile_modeld as stock
|
||||
from openpilot.selfdrive.modeld.compile_modeld import NV12Frame, make_frame_prepare, sample_desire, sample_skip, shift_and_sample
|
||||
from tinygrad import dtypes
|
||||
from tinygrad.device import Device
|
||||
from tinygrad.engine.jit import TinyJit
|
||||
@@ -41,7 +41,8 @@ from tinygrad.tensor import Tensor
|
||||
MODEL_TYPES = ('vision_policy', 'supercombo', 'vision_multi_policy')
|
||||
WARP_INPUTS = ['tfm', 'big_tfm']
|
||||
POLICY_INPUTS = ['img_q', 'big_img_q', 'feat_q', 'desire_q', 'packed_npy_inputs']
|
||||
nv12_copy_size = stock.nv12_copy_size
|
||||
WARP_DEV = os.getenv('WARP_DEV')
|
||||
|
||||
|
||||
def _detect_desire_key(shapes: dict) -> str | None:
|
||||
return next((key for key in shapes if key.startswith('desire')), None)
|
||||
@@ -138,7 +139,7 @@ def make_supercombo_input_queues(input_shapes: dict, frame_skip: int,
|
||||
return generate_queues_and_npy(input_shapes, frame_skip, device, is_supercombo=True)
|
||||
|
||||
|
||||
def make_random_images(keys, shape, device, rng=None):
|
||||
def make_random_images(keys, shape, device):
|
||||
return {k: Tensor.randint(shape, low=0, high=256, dtype=dtypes.uint8, device=device).realize() for k in keys}
|
||||
|
||||
|
||||
@@ -151,9 +152,24 @@ def make_warp_queues(device=Device.DEFAULT):
|
||||
return queues, npy
|
||||
|
||||
|
||||
def make_warp(nv12: NV12Frame, model_w: int, model_h: int):
|
||||
frame_prepare = make_frame_prepare(nv12, model_w, model_h)
|
||||
WARP_DEV = os.getenv('WARP_DEV', Device.DEFAULT)
|
||||
|
||||
def warp(tfm, big_tfm, frame, big_frame):
|
||||
tfm = tfm.to(WARP_DEV)
|
||||
big_tfm = big_tfm.to(WARP_DEV)
|
||||
Tensor.realize(tfm, big_tfm)
|
||||
|
||||
warped_frame = frame_prepare(frame, tfm).unsqueeze(0)
|
||||
warped_big_frame = frame_prepare(big_frame, big_tfm).unsqueeze(0)
|
||||
return Tensor.cat(warped_frame, warped_big_frame)
|
||||
return warp
|
||||
|
||||
|
||||
def make_run_policy(vision_runner, policy_runners: list, features_slice: slice, frame_skip: int, input_shapes: dict):
|
||||
sample_skip_fn = partial(stock.sample_skip, frame_skip=frame_skip)
|
||||
sample_desire_fn = partial(stock.sample_desire, frame_skip=frame_skip)
|
||||
sample_skip_fn = partial(sample_skip, frame_skip=frame_skip)
|
||||
sample_desire_fn = partial(sample_desire, frame_skip=frame_skip)
|
||||
|
||||
desire_key = _detect_desire_key(input_shapes)
|
||||
road_key, wide_key = _detect_vision_keys(input_shapes)
|
||||
@@ -170,14 +186,14 @@ def make_run_policy(vision_runner, policy_runners: list, features_slice: slice,
|
||||
warped_dev = warped.to(Device.DEFAULT)
|
||||
Tensor.realize(packed_npy_inputs_dev, warped_dev)
|
||||
|
||||
img = stock.shift_and_sample(img_q, warped_dev[0:1], sample_skip_fn)
|
||||
big_img = stock.shift_and_sample(big_img_q, warped_dev[1:2], sample_skip_fn)
|
||||
img = shift_and_sample(img_q, warped_dev[0:1], sample_skip_fn)
|
||||
big_img = shift_and_sample(big_img_q, warped_dev[1:2], sample_skip_fn)
|
||||
|
||||
unpacked_tensors = [tensor.reshape(shape) for tensor, shape in zip(packed_npy_inputs_dev.split(npy_sizes), npy_shapes.values(), strict=True)]
|
||||
unpacked_dict = dict(zip(npy_shapes.keys(), unpacked_tensors, strict=True))
|
||||
|
||||
desire_dev = unpacked_dict['desire']
|
||||
desire_buf = stock.shift_and_sample(desire_q, desire_dev.reshape(1, 1, -1), sample_desire_fn)
|
||||
desire_buf = shift_and_sample(desire_q, desire_dev.reshape(1, 1, -1), sample_desire_fn)
|
||||
|
||||
inputs = {desire_key: desire_buf}
|
||||
for key, tensor_val in unpacked_dict.items():
|
||||
@@ -186,13 +202,13 @@ def make_run_policy(vision_runner, policy_runners: list, features_slice: slice,
|
||||
|
||||
if 'prev_feat' in unpacked_dict:
|
||||
prev_feat_dev = unpacked_dict['prev_feat']
|
||||
inputs['features_buffer'] = stock.shift_and_sample(feat_q, prev_feat_dev.reshape(1, 1, -1), sample_skip_fn).reshape(input_shapes['features_buffer'])
|
||||
inputs['features_buffer'] = shift_and_sample(feat_q, prev_feat_dev.reshape(1, 1, -1), sample_skip_fn).reshape(input_shapes['features_buffer'])
|
||||
|
||||
if vision_runner:
|
||||
vision_out_cast = next(iter(vision_runner({road_key: img, wide_key: big_img}).values())).cast('float32').realize()
|
||||
if 'features_buffer' not in inputs:
|
||||
new_feat = vision_out_cast[:, features_slice].reshape(1, -1).unsqueeze(0)
|
||||
inputs['features_buffer'] = stock.shift_and_sample(feat_q, new_feat, sample_skip_fn).realize()
|
||||
inputs['features_buffer'] = shift_and_sample(feat_q, new_feat, sample_skip_fn).realize()
|
||||
policy_outs = [next(iter(pol_runner(inputs).values())).cast('float32').realize() for pol_runner in policy_runners]
|
||||
return (vision_out_cast, *policy_outs) if len(policy_outs) > 1 else (vision_out_cast, policy_outs[0])
|
||||
|
||||
@@ -203,28 +219,27 @@ def make_run_policy(vision_runner, policy_runners: list, features_slice: slice,
|
||||
policy_out = next(iter(policy_runners[0](inputs).values())).cast('float32').realize()
|
||||
if 'features_buffer' not in inputs and features_slice is not None:
|
||||
new_feat = policy_out[:, features_slice].reshape(1, -1).unsqueeze(0)
|
||||
stock.shift_and_sample(feat_q, new_feat, sample_skip_fn).realize()
|
||||
shift_and_sample(feat_q, new_feat, sample_skip_fn).realize()
|
||||
return policy_out
|
||||
|
||||
return run_policy
|
||||
|
||||
|
||||
def compile_jit(jit, input_keys, make_queues, make_random_inputs=None, benchmark_runs: int = 1):
|
||||
def compile_jit(jit, make_random_inputs, input_keys, make_queues):
|
||||
SEED = 42
|
||||
def random_inputs_run(fn, seed, n_runs, test_val=None, test_buffers=None, expect_match=True):
|
||||
queues_res = make_queues(Device.DEFAULT)
|
||||
input_queues, npy = queues_res[0], queues_res[1]
|
||||
frame_views = queues_res[2] if len(queues_res) > 2 else {}
|
||||
def random_inputs_run(fn, seed, test_val=None, test_buffers=None, expect_match=True):
|
||||
input_queues, npy = make_queues(Device.DEFAULT)
|
||||
rng = np.random.default_rng(seed)
|
||||
Tensor.manual_seed(seed)
|
||||
|
||||
testing = test_val is not None or test_buffers is not None
|
||||
n_runs = 1 if testing else 3
|
||||
|
||||
for i in range(n_runs):
|
||||
for v in npy.values():
|
||||
v[:] = rng.standard_normal(v.shape).astype(v.dtype)
|
||||
for v in frame_views.values():
|
||||
v[:] = rng.integers(0, 256, size=v.shape, dtype=np.uint8)
|
||||
Device.default.synchronize()
|
||||
random_inputs = make_random_inputs(rng=rng) if make_random_inputs is not None else {}
|
||||
random_inputs = make_random_inputs()
|
||||
st = time.perf_counter()
|
||||
outs = fn(**{k: input_queues[k] for k in input_keys if k in input_queues}, **random_inputs)
|
||||
mt = time.perf_counter()
|
||||
@@ -245,15 +260,14 @@ def compile_jit(jit, input_keys, make_queues, make_random_inputs=None, benchmark
|
||||
return val, buffers
|
||||
|
||||
print('capture + replay')
|
||||
test_val, test_buffers = random_inputs_run(jit, SEED, 3)
|
||||
print(f'pickle round trip ({benchmark_runs} runs per seed)')
|
||||
test_val, test_buffers = random_inputs_run(jit, SEED)
|
||||
print('pickle round trip')
|
||||
with tempfile.TemporaryFile(dir=".") as f:
|
||||
dump_oob(jit, f)
|
||||
f.seek(0)
|
||||
loaded_jit = load_oob(f)
|
||||
random_inputs_run(loaded_jit, SEED, benchmark_runs, test_val, test_buffers, expect_match=True)
|
||||
random_inputs_run(loaded_jit, SEED+1, benchmark_runs, test_val, test_buffers, expect_match=False)
|
||||
return jit
|
||||
deserialized_jit = load_oob(f)
|
||||
random_inputs_run(deserialized_jit, SEED, test_val=test_val, test_buffers=test_buffers)
|
||||
return deserialized_jit
|
||||
|
||||
|
||||
def _parse_size(size_str: str) -> tuple[int, int]:
|
||||
@@ -303,7 +317,6 @@ if __name__ == "__main__":
|
||||
parser.add_argument('--model-size', type=_parse_size, required=True, help='model input WxH')
|
||||
parser.add_argument('--camera-resolutions', type=_parse_size, nargs='+', required=True)
|
||||
parser.add_argument('--frame-skip', type=int, default=None, help='frame skip value (auto-derived if not provided)')
|
||||
parser.add_argument('--benchmark-runs', type=int, default=1, help='benchmark runs')
|
||||
parser.add_argument('--output', required=True)
|
||||
|
||||
parser.add_argument('--vision-onnx', help='vision ONNX (for split models)')
|
||||
@@ -322,64 +335,48 @@ if __name__ == "__main__":
|
||||
args.on_policy_onnx = read_file_chunked_to_disk(args.on_policy_onnx)
|
||||
args.supercombo_onnx = read_file_chunked_to_disk(args.supercombo_onnx)
|
||||
|
||||
if args.model_type == 'supercombo':
|
||||
vision_runner = OnnxRunner(args.vision_onnx) if args.vision_onnx else None
|
||||
|
||||
if args.model_type == 'vision_policy':
|
||||
assert vision_runner and args.policy_onnx
|
||||
policy_runners = [OnnxRunner(args.policy_onnx)]
|
||||
output_data['metadata'] = {'vision': make_metadata_dict(args.vision_onnx), 'policy': make_metadata_dict(args.policy_onnx)}
|
||||
elif args.model_type == 'supercombo':
|
||||
assert args.supercombo_onnx
|
||||
model_metadata = make_metadata_dict(args.supercombo_onnx)
|
||||
output_data['metadata'] = {'model': model_metadata, **model_metadata}
|
||||
output_data['input_devices'] = {'model': Device.DEFAULT}
|
||||
output_data['run_model'] = {}
|
||||
derived_frame_skip = args.frame_skip or derive_frame_skip({}, model_metadata['input_shapes'])
|
||||
model_runner = OnnxRunner(args.supercombo_onnx)
|
||||
run_policy = stock.make_run_policy(model_runner, model_metadata, derived_frame_skip)
|
||||
for cam_w, cam_h in args.camera_resolutions:
|
||||
print(f"Compiling unified run_model JIT for {cam_w}x{cam_h}...")
|
||||
nv12 = stock.NV12Frame(cam_w, cam_h, *get_nv12_info(cam_w, cam_h))
|
||||
frame_copy_size = stock.nv12_copy_size(nv12.stride, nv12.y_height, nv12.uv_height)
|
||||
make_model_queues = partial(stock.make_input_queues, model_metadata['input_shapes'], derived_frame_skip,
|
||||
frame_copy_size=frame_copy_size)
|
||||
warp = stock.make_warp(nv12, model_w, model_h)
|
||||
run_model_jit = TinyJit(stock.make_run_model(warp, run_policy, model_metadata, frame_copy_size), prune=True)
|
||||
output_data['run_model'][(cam_w, cam_h)] = compile_jit(run_model_jit, stock.MODELD_INPUTS, make_model_queues, benchmark_runs=args.benchmark_runs)
|
||||
else:
|
||||
vision_runner = OnnxRunner(args.vision_onnx) if args.vision_onnx else None
|
||||
if args.model_type == 'vision_policy':
|
||||
assert vision_runner and args.policy_onnx
|
||||
policy_runners = [OnnxRunner(args.policy_onnx)]
|
||||
output_data['metadata'] = {'vision': make_metadata_dict(args.vision_onnx), 'policy': make_metadata_dict(args.policy_onnx)}
|
||||
elif args.model_type == 'vision_multi_policy':
|
||||
assert vision_runner
|
||||
policy_runners, policy_names = _load_policy_runners(args)
|
||||
output_data['metadata'] = {'vision': make_metadata_dict(args.vision_onnx)}
|
||||
for name in policy_names:
|
||||
runner_arg = getattr(args, f"{name}_onnx")
|
||||
output_data['metadata'][name] = make_metadata_dict(runner_arg)
|
||||
policy_runners = [OnnxRunner(args.supercombo_onnx)]
|
||||
output_data['metadata'] = {'model': make_metadata_dict(args.supercombo_onnx)}
|
||||
elif args.model_type == 'vision_multi_policy':
|
||||
assert vision_runner
|
||||
policy_runners, policy_names = _load_policy_runners(args)
|
||||
output_data['metadata'] = {'vision': make_metadata_dict(args.vision_onnx)}
|
||||
for name in policy_names:
|
||||
runner_arg = getattr(args, f"{name}_onnx")
|
||||
output_data['metadata'][name] = make_metadata_dict(runner_arg)
|
||||
|
||||
policy_keys = [key for key in output_data['metadata'].keys() if key != 'vision']
|
||||
first_policy_meta = output_data['metadata'][policy_keys[0]] if policy_keys else {}
|
||||
vision_meta = output_data['metadata'].get('vision', {})
|
||||
policy_keys = [key for key in output_data['metadata'].keys() if key != 'vision']
|
||||
first_policy_meta = output_data['metadata'][policy_keys[0]] if policy_keys else {}
|
||||
vision_meta = output_data['metadata'].get('vision', {})
|
||||
|
||||
derived_frame_skip = args.frame_skip or derive_frame_skip(vision_meta.get('input_shapes', {}), first_policy_meta.get('input_shapes', {}))
|
||||
all_shapes = {key: value for meta in output_data['metadata'].values() for key, value in meta['input_shapes'].items()}
|
||||
feat_meta = output_data['metadata'].get('vision') or output_data['metadata'].get('policy')
|
||||
assert feat_meta is not None
|
||||
features_slice = feat_meta['output_slices']['hidden_state']
|
||||
derived_frame_skip = args.frame_skip or derive_frame_skip(vision_meta.get('input_shapes', {}), first_policy_meta.get('input_shapes', {}))
|
||||
all_shapes = {key: value for meta in output_data['metadata'].values() for key, value in meta['input_shapes'].items()}
|
||||
feat_meta = output_data['metadata'].get('vision') or output_data['metadata'].get('model') or output_data['metadata'].get('policy')
|
||||
assert feat_meta is not None
|
||||
features_slice = feat_meta['output_slices']['hidden_state']
|
||||
is_supercombo = vision_runner is None
|
||||
|
||||
print(f"Compiling run_policy JIT (model_size={model_w}x{model_h}, frame_skip={derived_frame_skip})...")
|
||||
run_policy_func = make_run_policy(vision_runner, policy_runners, features_slice, derived_frame_skip, all_shapes)
|
||||
run_policy_jit = TinyJit(run_policy_func, prune=True)
|
||||
make_policy_queues = partial(generate_queues_and_npy, all_shapes, derived_frame_skip, is_supercombo=False)
|
||||
make_random_model_inputs = partial(make_random_images, keys=['warped'], shape=(2, 6, model_h // 2, model_w // 2), device=Device.DEFAULT)
|
||||
output_data['run_policy'] = compile_jit(run_policy_jit, POLICY_INPUTS, make_policy_queues, make_random_inputs=make_random_model_inputs)
|
||||
print(f"Compiling run_policy JIT (model_size={model_w}x{model_h}, frame_skip={derived_frame_skip})...")
|
||||
run_policy_func = make_run_policy(vision_runner, policy_runners, features_slice, derived_frame_skip, all_shapes)
|
||||
run_policy_jit = TinyJit(run_policy_func, prune=True)
|
||||
make_policy_queues = partial(generate_queues_and_npy, all_shapes, derived_frame_skip, is_supercombo=is_supercombo)
|
||||
make_random_model_inputs = partial(make_random_images, keys=['warped'], shape=(2, 6, model_h // 2, model_w // 2), device=WARP_DEV)
|
||||
output_data['run_policy'] = compile_jit(run_policy_jit, make_random_model_inputs, POLICY_INPUTS, make_policy_queues)
|
||||
|
||||
for cam_w, cam_h in args.camera_resolutions:
|
||||
print(f"Compiling warp JIT for {cam_w}x{cam_h}...")
|
||||
nv12 = stock.NV12Frame(cam_w, cam_h, *get_nv12_info(cam_w, cam_h))
|
||||
frame_copy_size = stock.nv12_copy_size(nv12.stride, nv12.y_height, nv12.uv_height)
|
||||
make_random_warp_inputs = partial(make_random_images, keys=['frame', 'big_frame'], shape=frame_copy_size, device=Device.DEFAULT)
|
||||
warp = TinyJit(stock.make_warp(nv12, model_w, model_h), prune=True)
|
||||
output_data[(cam_w, cam_h)] = compile_jit(warp, WARP_INPUTS, make_warp_queues, make_random_inputs=make_random_warp_inputs)
|
||||
|
||||
output_data['metadata']['warp_dev'] = Device.DEFAULT
|
||||
for cam_w, cam_h in args.camera_resolutions:
|
||||
print(f"Compiling warp JIT for {cam_w}x{cam_h}...")
|
||||
nv12 = NV12Frame(cam_w, cam_h, *get_nv12_info(cam_w, cam_h))
|
||||
make_random_warp_inputs = partial(make_random_images, keys=['frame', 'big_frame'], shape=nv12.size, device=WARP_DEV)
|
||||
warp = TinyJit(make_warp(nv12, model_w, model_h), prune=True)
|
||||
output_data[(cam_w, cam_h)] = compile_jit(warp, make_random_warp_inputs, WARP_INPUTS, make_warp_queues)
|
||||
|
||||
with open(args.output, "wb") as file:
|
||||
dump_oob(output_data, file)
|
||||
|
||||
@@ -14,8 +14,6 @@ class ModelConstants:
|
||||
|
||||
# model inputs constants
|
||||
MODEL_FREQ = 20
|
||||
MODEL_RUN_FREQ = 20
|
||||
MODEL_CONTEXT_FREQ = 5
|
||||
FEATURE_LEN = 512
|
||||
FULL_HISTORY_BUFFER_LEN = 99
|
||||
DESIRE_LEN = 8
|
||||
@@ -37,7 +35,6 @@ class ModelConstants:
|
||||
LANE_LINES_WIDTH = 2
|
||||
ROAD_EDGES_WIDTH = 2
|
||||
PLAN_WIDTH = 15
|
||||
ACTION_WIDTH = 2
|
||||
DESIRE_PRED_WIDTH = 8
|
||||
LAT_PLANNER_SOLUTION_WIDTH = 4
|
||||
DESIRED_CURV_WIDTH = 1
|
||||
|
||||
@@ -1,9 +1,26 @@
|
||||
from openpilot.sunnypilot.modeld_v2.constants import Meta
|
||||
from openpilot.cereal import custom
|
||||
from openpilot.sunnypilot.modeld_v2.meta_20hz import Meta20hz
|
||||
from openpilot.sunnypilot.models.helpers import get_active_bundle
|
||||
|
||||
ModelBundle = custom.ModelManagerSP.ModelBundle
|
||||
|
||||
|
||||
def load_meta_constants():
|
||||
"""
|
||||
Determines and loads the appropriate meta model class based on the metadata provided. The function checks
|
||||
specific keys and conditions within the provided metadata dictionary to identify the corresponding meta
|
||||
model class to return.
|
||||
|
||||
:param model_metadata: Dictionary containing metadata about the model. It includes
|
||||
details such as input shapes, output slices, and other configurations for identifying
|
||||
metadata-dependent meta model classes.
|
||||
:type model_metadata: dict
|
||||
:return: The appropriate meta model class (Meta, MetaSimPose, or MetaTombRaider)
|
||||
based on the conditions and metadata provided.
|
||||
:rtype: type
|
||||
"""
|
||||
if (bundle := get_active_bundle()) and bundle.is20hz:
|
||||
return Meta20hz
|
||||
return Meta
|
||||
|
||||
return Meta # Default
|
||||
|
||||
@@ -6,7 +6,6 @@ This file is part of sunnypilot and is licensed under the MIT License.
|
||||
See the LICENSE.md file in the root directory for more details.
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
import os
|
||||
os.environ['GMMU'] = '0'
|
||||
import numpy as np
|
||||
@@ -38,18 +37,11 @@ from openpilot.selfdrive.controls.lib.desire_helper import DesireHelper
|
||||
from openpilot.selfdrive.controls.lib.drive_helpers import get_accel_from_plan, smooth_value
|
||||
from openpilot.selfdrive.modeld.modeld import ChestnutState
|
||||
|
||||
from openpilot.selfdrive.modeld.compile_modeld import (
|
||||
MODELD_INPUTS,
|
||||
make_input_queues as make_stock_input_queues,
|
||||
)
|
||||
from openpilot.sunnypilot.modeld_v2.fill_model_msg import fill_model_msg, fill_pose_msg, PublishState, get_curvature_from_output
|
||||
from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser
|
||||
from openpilot.sunnypilot.modeld_v2.constants import ModelConstants, Plan
|
||||
from openpilot.sunnypilot.modeld_v2.constants import Plan
|
||||
from openpilot.sunnypilot.modeld_v2.meta_helper import load_meta_constants
|
||||
from openpilot.sunnypilot.modeld_v2.camera_offset_helper import CameraOffsetHelper
|
||||
from openpilot.sunnypilot.modeld_v2.compile_modeld import (derive_frame_skip, make_split_input_queues,
|
||||
make_supercombo_input_queues, nv12_copy_size,
|
||||
WARP_INPUTS, POLICY_INPUTS)
|
||||
from openpilot.sunnypilot.modeld_v2.compile_modeld import derive_frame_skip, make_split_input_queues, make_supercombo_input_queues, WARP_INPUTS, POLICY_INPUTS
|
||||
from openpilot.sunnypilot.livedelay.helpers import get_lat_delay
|
||||
from openpilot.sunnypilot.modeld_v2.modeld_base import ModelStateBase
|
||||
from openpilot.sunnypilot.models.helpers import get_active_bundle
|
||||
@@ -118,40 +110,36 @@ class ModelState(ModelStateBase):
|
||||
cloudlog.warning(f"loading combined pkl: {pkl_path}")
|
||||
jits = load_oob(open_file_chunked(pkl_path))
|
||||
|
||||
metadata = jits['metadata']
|
||||
self.WARP_DEV = metadata.get('warp_dev', 'QCOM') if COMMA_HARDWARE else 'CPU'
|
||||
self.DEV = ('AMD' if self.chestnut else 'QCOM') if COMMA_HARDWARE else 'CPU'
|
||||
self.WARP_DEV = 'QCOM' if COMMA_HARDWARE else 'CPU'
|
||||
self.DEV = 'AMD' if self.chestnut else self.WARP_DEV
|
||||
self.QUEUE_DEV = self.DEV
|
||||
self.is_run_model = 'run_model' in jits
|
||||
metadata = jits['metadata']
|
||||
|
||||
nv12_info = get_nv12_info(cam_w, cam_h)
|
||||
self.frame_copy_size = nv12_copy_size(*nv12_info[:3])
|
||||
self.full_frames: dict = {}
|
||||
self._blob_cache: dict = {}
|
||||
self.frame_buffers: dict = {}
|
||||
self.is_legacy_model = 'run_policy' not in jits # remove after next recompile
|
||||
if self.is_legacy_model:
|
||||
self.warp = jits[(cam_w, cam_h)]['warp_enqueue']
|
||||
self.run_policy = jits[(cam_w, cam_h)]['run_policy']
|
||||
else:
|
||||
self.run_policy = jits['run_policy']
|
||||
self.warp = jits[(cam_w, cam_h)]
|
||||
|
||||
if self.is_run_model or 'model' in metadata:
|
||||
model_metadata = metadata.get('model', metadata)
|
||||
self.input_shapes = model_metadata['input_shapes']
|
||||
if 'model' in metadata:
|
||||
model_metadata = metadata['model']
|
||||
self.vision_output_slices = model_metadata['output_slices']
|
||||
self.policy_output_slices = {}
|
||||
self._policy_slices_list = []
|
||||
self._combined_model_type = 'supercombo'
|
||||
self._vision_input_names = [key for key in self.input_shapes if 'img' in key]
|
||||
self.frame_skip = derive_frame_skip({}, self.input_shapes)
|
||||
if self.is_run_model:
|
||||
self.input_queues, self.numpy_inputs, self.frame_buffers = make_stock_input_queues(
|
||||
self.input_shapes, self.frame_skip, device=self.DEV, frame_copy_size=self.frame_copy_size)
|
||||
self.frame_views, self.npy = self.frame_buffers, self.numpy_inputs
|
||||
self.run_model, self.run_policy, self.warp = jits['run_model'][(cam_w, cam_h)], None, None
|
||||
else:
|
||||
self.input_queues, self.numpy_inputs = make_supercombo_input_queues(self.input_shapes, self.frame_skip, device=self.QUEUE_DEV)
|
||||
self.run_model, self.run_policy, self.warp = None, jits['run_policy'], jits[(cam_w, cam_h)]
|
||||
self._vision_input_names = [key for key in model_metadata['input_shapes'] if 'img' in key]
|
||||
frame_skip = derive_frame_skip({}, model_metadata['input_shapes'])
|
||||
self.input_queues, self.numpy_inputs = make_supercombo_input_queues(model_metadata['input_shapes'],
|
||||
frame_skip, device=self.QUEUE_DEV)
|
||||
else:
|
||||
self.run_model, self.run_policy, self.warp = None, jits['run_policy'], jits[(cam_w, cam_h)]
|
||||
vision_metadata = metadata['vision']
|
||||
policy_keys = [k for k in metadata if k not in ('vision', 'warp_dev')]
|
||||
self._combined_model_type = 'split' if policy_keys == ['policy'] else 'multi_policy'
|
||||
policy_keys = [k for k in metadata if k != 'vision']
|
||||
if policy_keys == ['policy']:
|
||||
self._combined_model_type = 'split'
|
||||
else:
|
||||
self._combined_model_type = 'multi_policy'
|
||||
self.vision_output_slices = vision_metadata['output_slices']
|
||||
self._policy_keys = policy_keys
|
||||
self._policy_slices_list = [metadata[k]['output_slices'] for k in policy_keys]
|
||||
@@ -167,39 +155,54 @@ class ModelState(ModelStateBase):
|
||||
self._desire_key = next(key for key in self.numpy_inputs if key.startswith('desire'))
|
||||
self._road_key = next(key for key in self._vision_input_names if 'big' not in key)
|
||||
self._wide_key = next(key for key in self._vision_input_names if 'big' in key)
|
||||
self.frame_buf_params = dict.fromkeys(self._vision_input_names, nv12_info)
|
||||
|
||||
is_20hz = bundle.is20hz if bundle else self._combined_model_type in ('split', 'multi_policy')
|
||||
if is_20hz:
|
||||
from openpilot.sunnypilot.models.split_model_constants import SplitModelConstants
|
||||
self.constants = SplitModelConstants()
|
||||
else:
|
||||
from openpilot.sunnypilot.modeld_v2.constants import ModelConstants
|
||||
self.constants = ModelConstants()
|
||||
|
||||
self.parser = Parser()
|
||||
self.prev_desire = np.zeros(self.constants.DESIRE_LEN, dtype=np.float32)
|
||||
if self._combined_model_type != 'supercombo':
|
||||
from openpilot.sunnypilot.modeld_v2.parse_model_outputs_split import Parser as SplitParser
|
||||
self.parser = SplitParser()
|
||||
else:
|
||||
from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser as CombinedParser
|
||||
self.parser = CombinedParser()
|
||||
|
||||
if self.warp is not None:
|
||||
self.full_frames = {k: Tensor(np.zeros(nv12_info[3], dtype=np.uint8), device=self.WARP_DEV).contiguous().realize() for k in self._vision_input_names}
|
||||
self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[self._road_key], big_frame=self.full_frames[self._wide_key])
|
||||
self.prev_desire = np.zeros(self.constants.DESIRE_LEN, dtype=np.float32)
|
||||
self.full_frames: dict = {}
|
||||
self._blob_cache: dict = {}
|
||||
nv12_info = get_nv12_info(cam_w, cam_h)
|
||||
self.frame_buf_params = dict.fromkeys(self._vision_input_names, nv12_info)
|
||||
|
||||
yuv_size = self.frame_buf_params[self._road_key][3]
|
||||
frame_tensor = Tensor(np.zeros(yuv_size, dtype=np.uint8), device=self.WARP_DEV).contiguous().realize()
|
||||
big_frame_tensor = Tensor(np.zeros(yuv_size, dtype=np.uint8), device=self.WARP_DEV).contiguous().realize()
|
||||
|
||||
if self.is_legacy_model: # Remove this conditional hack after recompile
|
||||
self.warp(**self.input_queues, frame=frame_tensor, big_frame=big_frame_tensor)
|
||||
else:
|
||||
self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=frame_tensor, big_frame=big_frame_tensor)
|
||||
|
||||
def warmup(self) -> None:
|
||||
dummy_size = self.frame_copy_size if self.is_run_model else self.frame_buf_params[self._road_key][3]
|
||||
dummy_frames = {k: np.zeros(dummy_size, dtype=np.uint8) for k in self._vision_input_names}
|
||||
dummy_frames = {k: np.zeros(self.frame_buf_params[k][3], dtype=np.uint8) for k in self._vision_input_names}
|
||||
transforms = {k: np.eye(3, dtype=np.float32) for k in [self._road_key, self._wide_key] if k}
|
||||
dummy_inputs = {k: np.zeros(v.shape, dtype=v.dtype) for k, v in self.numpy_inputs.items() if k not in ['tfm', 'big_tfm', 'prev_feat']}
|
||||
self.run(dummy_frames, transforms, dummy_inputs)
|
||||
if self.is_run_model:
|
||||
self.input_queues, self.numpy_inputs, self.frame_buffers = make_stock_input_queues(
|
||||
self.input_shapes, self.frame_skip, device=self.DEV, frame_copy_size=self.frame_copy_size)
|
||||
self.frame_views = self.frame_buffers
|
||||
self.npy = self.numpy_inputs
|
||||
else:
|
||||
for v in self.numpy_inputs.values():
|
||||
v[:] = 0
|
||||
self.full_frames.clear()
|
||||
self._blob_cache.clear()
|
||||
|
||||
dummy_inputs = {}
|
||||
for k, v in self.numpy_inputs.items():
|
||||
if k not in ['tfm', 'big_tfm', 'prev_feat']:
|
||||
dummy_inputs[k] = np.zeros(v.shape, dtype=v.dtype)
|
||||
|
||||
self.run(dummy_frames, transforms, dummy_inputs, prepare_only=False)
|
||||
|
||||
for v in self.numpy_inputs.values():
|
||||
v[:] = 0
|
||||
self.prev_desire[:] = 0
|
||||
self.full_frames.clear()
|
||||
self._blob_cache.clear()
|
||||
|
||||
|
||||
@property
|
||||
def mlsim(self) -> bool:
|
||||
@@ -214,50 +217,45 @@ class ModelState(ModelStateBase):
|
||||
return self._desire_key
|
||||
|
||||
def run(self, bufs: dict[str, VisionBuf], transforms: dict[str, np.ndarray],
|
||||
inputs: dict[str, np.ndarray],
|
||||
after_enqueue: Callable[[], None] | None = None) -> dict[str, np.ndarray] | None:
|
||||
if self.is_run_model:
|
||||
for key, buf in bufs.items():
|
||||
data = buf.data if hasattr(buf, 'data') else buf
|
||||
np.copyto(self.frame_buffers[key], np.frombuffer(data, dtype=np.uint8, count=self.frame_copy_size))
|
||||
else:
|
||||
for key, buf in bufs.items():
|
||||
ptr = np.frombuffer(buf.data, dtype=np.uint8).ctypes.data
|
||||
cache_key = (key, ptr)
|
||||
if cache_key not in self._blob_cache:
|
||||
self._blob_cache[cache_key] = Tensor.from_blob(ptr, (self.frame_buf_params[key][3],), dtype='uint8', device=self.WARP_DEV)
|
||||
self.full_frames[key] = self._blob_cache[cache_key]
|
||||
inputs: dict[str, np.ndarray], prepare_only: bool) -> dict[str, np.ndarray] | None:
|
||||
for key in bufs.keys():
|
||||
ptr = np.frombuffer(bufs[key].data, dtype=np.uint8).ctypes.data
|
||||
yuv_size = self.frame_buf_params[key][3]
|
||||
cache_key = (key, ptr)
|
||||
if cache_key not in self._blob_cache:
|
||||
self._blob_cache[cache_key] = Tensor.from_blob(ptr, (yuv_size,), dtype='uint8', device=self.WARP_DEV)
|
||||
self.full_frames[key] = self._blob_cache[cache_key]
|
||||
|
||||
desire_key = self.desire_key
|
||||
inputs[desire_key][0] = 0
|
||||
self.numpy_inputs[desire_key][:] = np.where(inputs[desire_key] - self.prev_desire > .99, inputs[desire_key], 0)
|
||||
self.prev_desire[:] = inputs[desire_key]
|
||||
|
||||
for key in ('traffic_convention', 'lateral_control_params', 'action_t'):
|
||||
if key in self.numpy_inputs and key in inputs:
|
||||
self.numpy_inputs[key][:] = inputs[key]
|
||||
|
||||
self.numpy_inputs['tfm'][:, :] = transforms[self._road_key].reshape(3, 3)
|
||||
self.numpy_inputs['big_tfm'][:, :] = transforms[self._wide_key].reshape(3, 3)
|
||||
road_key = self._road_key
|
||||
wide_key = self._wide_key
|
||||
self.numpy_inputs['tfm'][:, :] = transforms[road_key].reshape(3, 3)
|
||||
self.numpy_inputs['big_tfm'][:, :] = transforms[wide_key].reshape(3, 3)
|
||||
|
||||
if self.run_model is not None:
|
||||
outs, = self.run_model(**{k: self.input_queues[k] for k in MODELD_INPUTS})
|
||||
raw_outputs = outs
|
||||
if self.is_legacy_model: # remove after next recompile
|
||||
if prepare_only:
|
||||
self.warp(**self.input_queues, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key])
|
||||
return None
|
||||
raw_outputs = self.run_policy(**self.input_queues, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key])
|
||||
else:
|
||||
assert self.warp is not None and self.run_policy is not None
|
||||
warped = self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[self._road_key], big_frame=self.full_frames[self._wide_key])
|
||||
if prepare_only:
|
||||
self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key])
|
||||
return None
|
||||
warped = self.warp(**{k: self.input_queues[k] for k in WARP_INPUTS}, frame=self.full_frames[road_key], big_frame=self.full_frames[wide_key])
|
||||
raw_outputs = self.run_policy(**{k: self.input_queues[k] for k in POLICY_INPUTS if k in self.input_queues}, warped=warped)
|
||||
|
||||
if after_enqueue is not None:
|
||||
after_enqueue()
|
||||
|
||||
if self._combined_model_type == 'supercombo':
|
||||
model_output = raw_outputs.numpy().flatten()
|
||||
if self.chestnut and not np.all(np.isfinite(model_output)):
|
||||
raise RuntimeError("model output not finite")
|
||||
sliced = {k: model_output[np.newaxis, v] for k, v in self.vision_output_slices.items()}
|
||||
outputs = self.parser.parse_outputs(sliced)
|
||||
if 'prev_feat' in self.numpy_inputs and 'hidden_state' in self.vision_output_slices:
|
||||
if 'prev_feat' in self.numpy_inputs:
|
||||
self.numpy_inputs['prev_feat'][:] = model_output[self.vision_output_slices['hidden_state']]
|
||||
else:
|
||||
vision_output = raw_outputs[0].numpy().flatten()
|
||||
@@ -287,6 +285,9 @@ class ModelState(ModelStateBase):
|
||||
buf[0, :-1] = buf[0, 1:]
|
||||
buf[0, -1, :] = outputs['desired_curvature'][0, :] if not self.mlsim else 0
|
||||
|
||||
if self.chestnut and not np.all(np.isfinite(outputs.get('plan', np.array([0.])))):
|
||||
raise RuntimeError("model output not finite")
|
||||
|
||||
return outputs
|
||||
|
||||
def get_action_from_model(self, model_output: dict[str, np.ndarray], prev_action: log.ModelDataV2.Action,
|
||||
@@ -372,11 +373,7 @@ def main(demo=False):
|
||||
loader.start()
|
||||
loader.join(BIG_MODEL_TIMEOUT)
|
||||
model = big_model
|
||||
if model is None:
|
||||
params.put_bool("ChestnutModelError", True)
|
||||
params.put_bool("ChestnutActive", model is not None)
|
||||
if model is not None:
|
||||
params.remove("ChestnutModelError")
|
||||
|
||||
small_model = ModelState(cam_w=vipc_client_main.width, cam_h=vipc_client_main.height, chestnut=False) if model is None or CHESTNUT else None
|
||||
if model is None:
|
||||
@@ -490,6 +487,9 @@ def main(demo=False):
|
||||
run_count = run_count + 1
|
||||
|
||||
frame_drop_ratio = frames_dropped / (1 + frames_dropped)
|
||||
prepare_only = vipc_dropped_frames > 0
|
||||
if prepare_only:
|
||||
cloudlog.error(f"skipping model eval. Dropped {vipc_dropped_frames} frames")
|
||||
|
||||
bufs = {name: buf_extra if 'big' in name else buf_main for name in model.vision_input_names}
|
||||
transforms = {name: model_transform_extra if 'big' in name else model_transform_main for name in model.vision_input_names}
|
||||
@@ -512,14 +512,11 @@ def main(demo=False):
|
||||
|
||||
mt1 = time.perf_counter()
|
||||
try:
|
||||
send_chestnut = (chestnut_state is not None and
|
||||
run_count % round(model.constants.MODEL_FREQ / SERVICE_LIST['chestnutState'].frequency) == 0)
|
||||
model_output = model.run(bufs, transforms, inputs, chestnut_state.send if send_chestnut else None)
|
||||
model_output = model.run(bufs, transforms, inputs, prepare_only)
|
||||
except Exception:
|
||||
if not params.get_bool("ChestnutActive"):
|
||||
raise
|
||||
cloudlog.exception("chestnut failed, falling back to small")
|
||||
params.put_bool("ChestnutModelError", True)
|
||||
params.put_bool("ChestnutActive", False)
|
||||
assert small_model is not None
|
||||
model = small_model
|
||||
@@ -562,6 +559,9 @@ def main(demo=False):
|
||||
pm.send('modelDataV2SP', mdv2sp_send)
|
||||
last_vipc_frame_id = meta_main.frame_id
|
||||
|
||||
if chestnut_state is not None and run_count % round(model.constants.MODEL_FREQ / SERVICE_LIST['chestnutState'].frequency) == 0:
|
||||
chestnut_state.send()
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
import argparse
|
||||
|
||||
@@ -115,41 +115,22 @@ class Parser:
|
||||
outs[name + '_stds'] = pred_std_final.reshape(final_shape)
|
||||
|
||||
def parse_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
if 'plan' in outs:
|
||||
self.parse_mdn('plan', outs, out_shape=(ModelConstants.IDX_N, ModelConstants.PLAN_WIDTH))
|
||||
if 'planplus' in outs:
|
||||
self.parse_mdn('planplus', outs, out_shape=(ModelConstants.IDX_N, ModelConstants.PLAN_WIDTH))
|
||||
if 'lane_lines' in outs:
|
||||
self.parse_mdn('lane_lines', outs, out_shape=(ModelConstants.NUM_LANE_LINES, ModelConstants.IDX_N, ModelConstants.LANE_LINES_WIDTH))
|
||||
if 'road_edges' in outs:
|
||||
self.parse_mdn('road_edges', outs, out_shape=(ModelConstants.NUM_ROAD_EDGES, ModelConstants.IDX_N, ModelConstants.LANE_LINES_WIDTH))
|
||||
if 'pose' in outs:
|
||||
self.parse_mdn('pose', outs, out_shape=(ModelConstants.POSE_WIDTH,))
|
||||
if 'road_transform' in outs:
|
||||
self.parse_mdn('road_transform', outs, out_shape=(ModelConstants.POSE_WIDTH,))
|
||||
# supercombo (4955 / 102) and newer variants (e.g. 990 / 144).
|
||||
self.parse_mdn('plan', outs, out_shape=(ModelConstants.IDX_N, ModelConstants.PLAN_WIDTH))
|
||||
self.parse_mdn('lane_lines', outs, out_shape=(ModelConstants.NUM_LANE_LINES, ModelConstants.IDX_N, ModelConstants.LANE_LINES_WIDTH))
|
||||
self.parse_mdn('road_edges', outs, out_shape=(ModelConstants.NUM_ROAD_EDGES, ModelConstants.IDX_N, ModelConstants.LANE_LINES_WIDTH))
|
||||
self.parse_mdn('pose', outs, out_shape=(ModelConstants.POSE_WIDTH,))
|
||||
self.parse_mdn('road_transform', outs, out_shape=(ModelConstants.POSE_WIDTH,))
|
||||
if 'sim_pose' in outs:
|
||||
self.parse_mdn('sim_pose', outs, out_shape=(ModelConstants.POSE_WIDTH,))
|
||||
if 'wide_from_device_euler' in outs:
|
||||
self.parse_mdn('wide_from_device_euler', outs, out_shape=(ModelConstants.WIDE_FROM_DEVICE_WIDTH,))
|
||||
if 'lead' in outs:
|
||||
self.parse_mdn('lead', outs, out_shape=(ModelConstants.LEAD_TRAJ_LEN, ModelConstants.LEAD_WIDTH))
|
||||
self.parse_mdn('wide_from_device_euler', outs, out_shape=(ModelConstants.WIDE_FROM_DEVICE_WIDTH,))
|
||||
self.parse_mdn('lead', outs, out_shape=(ModelConstants.LEAD_TRAJ_LEN, ModelConstants.LEAD_WIDTH))
|
||||
if 'lat_planner_solution' in outs:
|
||||
self.parse_mdn('lat_planner_solution', outs, out_shape=(ModelConstants.IDX_N, ModelConstants.LAT_PLANNER_SOLUTION_WIDTH))
|
||||
if 'desired_curvature' in outs:
|
||||
self.parse_mdn('desired_curvature', outs, out_shape=(ModelConstants.DESIRED_CURV_WIDTH,))
|
||||
if 'action' in outs:
|
||||
self.parse_mdn('action', outs, out_shape=(ModelConstants.ACTION_WIDTH,))
|
||||
for k in ['lead_prob', 'lane_lines_prob', 'meta']:
|
||||
if k in outs:
|
||||
self.parse_binary_crossentropy(k, outs)
|
||||
if 'desire_state' in outs:
|
||||
self.parse_categorical_crossentropy('desire_state', outs, out_shape=(ModelConstants.DESIRE_PRED_WIDTH,))
|
||||
if 'desire_pred' in outs:
|
||||
self.parse_categorical_crossentropy('desire_pred', outs, out_shape=(ModelConstants.DESIRE_PRED_LEN, ModelConstants.DESIRE_PRED_WIDTH))
|
||||
self.parse_binary_crossentropy(k, outs)
|
||||
self.parse_categorical_crossentropy('desire_state', outs, out_shape=(ModelConstants.DESIRE_PRED_WIDTH,))
|
||||
self.parse_categorical_crossentropy('desire_pred', outs, out_shape=(ModelConstants.DESIRE_PRED_LEN, ModelConstants.DESIRE_PRED_WIDTH))
|
||||
return outs
|
||||
|
||||
def parse_vision_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
return self.parse_outputs(outs)
|
||||
|
||||
def parse_policy_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
return self.parse_outputs(outs)
|
||||
|
||||
@@ -0,0 +1,159 @@
|
||||
import numpy as np
|
||||
from openpilot.sunnypilot.models.split_model_constants import SplitModelConstants
|
||||
|
||||
|
||||
def safe_exp(x, out=None):
|
||||
# -11 is around 10**14, more causes float16 overflow
|
||||
return np.exp(np.clip(x, -np.inf, 11), out=out)
|
||||
|
||||
|
||||
def sigmoid(x):
|
||||
return 1. / (1. + safe_exp(-x))
|
||||
|
||||
|
||||
def softmax(x, axis=-1):
|
||||
x -= np.max(x, axis=axis, keepdims=True)
|
||||
if x.dtype == np.float32 or x.dtype == np.float64:
|
||||
safe_exp(x, out=x)
|
||||
else:
|
||||
x = safe_exp(x)
|
||||
x /= np.sum(x, axis=axis, keepdims=True)
|
||||
return x
|
||||
|
||||
|
||||
class Parser:
|
||||
def __init__(self, ignore_missing=False):
|
||||
self.ignore_missing = ignore_missing
|
||||
|
||||
def check_missing(self, outs, name):
|
||||
if name not in outs and not self.ignore_missing:
|
||||
raise ValueError(f"Missing output {name}")
|
||||
return name not in outs
|
||||
|
||||
def parse_categorical_crossentropy(self, name, outs, out_shape=None):
|
||||
if self.check_missing(outs, name):
|
||||
return
|
||||
raw = outs[name]
|
||||
if out_shape is not None:
|
||||
raw = raw.reshape((raw.shape[0],) + out_shape)
|
||||
outs[name] = softmax(raw, axis=-1)
|
||||
|
||||
def parse_binary_crossentropy(self, name, outs):
|
||||
if self.check_missing(outs, name):
|
||||
return
|
||||
raw = outs[name]
|
||||
outs[name] = sigmoid(raw)
|
||||
|
||||
def parse_mdn(self, name, outs, in_N=0, out_N=1, out_shape=None):
|
||||
if self.check_missing(outs, name):
|
||||
return
|
||||
raw = outs[name]
|
||||
raw = raw.reshape((raw.shape[0], max(in_N, 1), -1))
|
||||
|
||||
n_values = (raw.shape[2] - out_N)//2
|
||||
pred_mu = raw[:,:,:n_values]
|
||||
pred_std = safe_exp(raw[:,:,n_values: 2*n_values])
|
||||
|
||||
if in_N > 1:
|
||||
weights = np.zeros((raw.shape[0], in_N, out_N), dtype=raw.dtype)
|
||||
for i in range(out_N):
|
||||
weights[:,:,i - out_N] = softmax(raw[:,:,i - out_N], axis=-1)
|
||||
|
||||
if out_N == 1:
|
||||
for fidx in range(weights.shape[0]):
|
||||
idxs = np.argsort(weights[fidx][:,0])[::-1]
|
||||
weights[fidx] = weights[fidx][idxs]
|
||||
pred_mu[fidx] = pred_mu[fidx][idxs]
|
||||
pred_std[fidx] = pred_std[fidx][idxs]
|
||||
assert out_shape is not None
|
||||
full_shape = tuple([raw.shape[0], in_N] + list(out_shape))
|
||||
outs[name + '_weights'] = weights
|
||||
outs[name + '_hypotheses'] = pred_mu.reshape(full_shape)
|
||||
outs[name + '_stds_hypotheses'] = pred_std.reshape(full_shape)
|
||||
|
||||
pred_mu_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype)
|
||||
pred_std_final = np.zeros((raw.shape[0], out_N, n_values), dtype=raw.dtype)
|
||||
for fidx in range(weights.shape[0]):
|
||||
for hidx in range(out_N):
|
||||
idxs = np.argsort(weights[fidx,:,hidx])[::-1]
|
||||
pred_mu_final[fidx, hidx] = pred_mu[fidx, idxs[0]]
|
||||
pred_std_final[fidx, hidx] = pred_std[fidx, idxs[0]]
|
||||
else:
|
||||
pred_mu_final = pred_mu
|
||||
pred_std_final = pred_std
|
||||
|
||||
if out_N > 1:
|
||||
assert out_shape is not None
|
||||
final_shape = tuple([raw.shape[0], out_N] + list(out_shape))
|
||||
else:
|
||||
assert out_shape is not None
|
||||
final_shape = tuple([raw.shape[0],] + list(out_shape))
|
||||
outs[name] = pred_mu_final.reshape(final_shape)
|
||||
outs[name + '_stds'] = pred_std_final.reshape(final_shape)
|
||||
|
||||
def is_mhp(self, outs, name, shape):
|
||||
if self.check_missing(outs, name):
|
||||
return False
|
||||
if outs[name].shape[1] == 2 * shape:
|
||||
return False
|
||||
return True
|
||||
|
||||
def parse_dynamic_outputs(self, outs: dict[str, np.ndarray]) -> None:
|
||||
if 'lead' in outs:
|
||||
lead_mhp = self.is_mhp(outs, 'lead',
|
||||
SplitModelConstants.LEAD_MHP_SELECTION * SplitModelConstants.LEAD_TRAJ_LEN * SplitModelConstants.LEAD_WIDTH)
|
||||
lead_in_N, lead_out_N = (SplitModelConstants.LEAD_MHP_N, SplitModelConstants.LEAD_MHP_SELECTION) if lead_mhp else (0, 0)
|
||||
lead_out_shape = (SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH) if lead_mhp else \
|
||||
(SplitModelConstants.LEAD_MHP_SELECTION, SplitModelConstants.LEAD_TRAJ_LEN, SplitModelConstants.LEAD_WIDTH)
|
||||
self.parse_mdn('lead', outs, in_N=lead_in_N, out_N=lead_out_N, out_shape=lead_out_shape)
|
||||
if 'plan' in outs:
|
||||
plan_mhp = self.is_mhp(outs, 'plan', SplitModelConstants.IDX_N * SplitModelConstants.PLAN_WIDTH)
|
||||
plan_in_N, plan_out_N = (SplitModelConstants.PLAN_MHP_N, SplitModelConstants.PLAN_MHP_SELECTION) if plan_mhp else (0, 0)
|
||||
self.parse_mdn('plan', outs, in_N=plan_in_N, out_N=plan_out_N,
|
||||
out_shape=(SplitModelConstants.IDX_N, SplitModelConstants.PLAN_WIDTH))
|
||||
if 'planplus' in outs:
|
||||
self.parse_mdn('planplus', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.IDX_N, SplitModelConstants.PLAN_WIDTH))
|
||||
|
||||
def split_outputs(self, outs: dict[str, np.ndarray]) -> None:
|
||||
if 'desired_curvature' in outs:
|
||||
self.parse_mdn('desired_curvature', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.DESIRED_CURV_WIDTH,))
|
||||
if 'desire_pred' in outs:
|
||||
self.parse_categorical_crossentropy('desire_pred', outs, out_shape=(SplitModelConstants.DESIRE_PRED_LEN,SplitModelConstants.DESIRE_PRED_WIDTH))
|
||||
if 'desire_state' in outs:
|
||||
self.parse_categorical_crossentropy('desire_state', outs, out_shape=(SplitModelConstants.DESIRE_PRED_WIDTH,))
|
||||
if 'lane_lines' in outs:
|
||||
self.parse_mdn('lane_lines', outs, in_N=0, out_N=0,
|
||||
out_shape=(SplitModelConstants.NUM_LANE_LINES,SplitModelConstants.IDX_N,SplitModelConstants.LANE_LINES_WIDTH))
|
||||
if 'lane_lines_prob' in outs:
|
||||
self.parse_binary_crossentropy('lane_lines_prob', outs)
|
||||
if 'lead_prob' in outs:
|
||||
self.parse_binary_crossentropy('lead_prob', outs)
|
||||
if 'lat_planner_solution' in outs:
|
||||
self.parse_mdn('lat_planner_solution', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.IDX_N,SplitModelConstants.LAT_PLANNER_SOLUTION_WIDTH))
|
||||
if 'meta' in outs:
|
||||
self.parse_binary_crossentropy('meta', outs)
|
||||
if 'road_edges' in outs:
|
||||
self.parse_mdn('road_edges', outs, in_N=0, out_N=0,
|
||||
out_shape=(SplitModelConstants.NUM_ROAD_EDGES,SplitModelConstants.IDX_N,SplitModelConstants.LANE_LINES_WIDTH))
|
||||
if 'sim_pose' in outs:
|
||||
self.parse_mdn('sim_pose', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.POSE_WIDTH,))
|
||||
if 'action' in outs:
|
||||
self.parse_mdn('action', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.ACTION_WIDTH,))
|
||||
|
||||
def parse_vision_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
self.parse_mdn('pose', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.POSE_WIDTH,))
|
||||
self.parse_mdn('wide_from_device_euler', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.WIDE_FROM_DEVICE_WIDTH,))
|
||||
self.parse_mdn('road_transform', outs, in_N=0, out_N=0, out_shape=(SplitModelConstants.POSE_WIDTH,))
|
||||
self.parse_dynamic_outputs(outs)
|
||||
self.split_outputs(outs)
|
||||
return outs
|
||||
|
||||
def parse_policy_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
self.parse_dynamic_outputs(outs)
|
||||
self.split_outputs(outs)
|
||||
return outs
|
||||
|
||||
def parse_outputs(self, outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
|
||||
outs = self.parse_vision_outputs(outs)
|
||||
outs = self.parse_policy_outputs(outs)
|
||||
return outs
|
||||
@@ -117,7 +117,7 @@ ARCHETYPES = {
|
||||
is_20hz=True,
|
||||
expected_model_type='split',
|
||||
expected_constants_class=SplitModelConstants,
|
||||
expected_parser_module='parse_model_outputs',
|
||||
expected_parser_module='parse_model_outputs_split',
|
||||
expected_desire_key='desire',
|
||||
),
|
||||
'vision_multi_policy': Archetype(
|
||||
@@ -130,7 +130,7 @@ ARCHETYPES = {
|
||||
is_20hz=True,
|
||||
expected_model_type='multi_policy',
|
||||
expected_constants_class=SplitModelConstants,
|
||||
expected_parser_module='parse_model_outputs',
|
||||
expected_parser_module='parse_model_outputs_split',
|
||||
expected_desire_key='desire',
|
||||
),
|
||||
'tri_policy': Archetype(
|
||||
@@ -144,7 +144,7 @@ ARCHETYPES = {
|
||||
is_20hz=True,
|
||||
expected_model_type='multi_policy',
|
||||
expected_constants_class=SplitModelConstants,
|
||||
expected_parser_module='parse_model_outputs',
|
||||
expected_parser_module='parse_model_outputs_split',
|
||||
expected_desire_key='desire',
|
||||
),
|
||||
'supercombo_non20hz': Archetype(
|
||||
|
||||
@@ -103,23 +103,6 @@ class TestStockEquivalence(OpenpilotTestCase):
|
||||
assert state.vision_output_slices == arch.metadata_structure['vision']['output_slices']
|
||||
assert state.policy_output_slices == arch.metadata_structure['policy']['output_slices']
|
||||
|
||||
def test_unified_run_model(self, tmp_path, monkeypatch, patch_modeld):
|
||||
from openpilot.common.hardware import hw
|
||||
from openpilot.selfdrive.modeld.helpers import dump_oob
|
||||
shapes = {'img': (1, 12, 128, 256), 'big_img': (1, 12, 128, 256), 'features_buffer': (1, 24, 32, 512),
|
||||
'desire_pulse': (1, 25, 8), 'traffic_convention': (1, 2), 'action_t': (1, 2)}
|
||||
pkl_data = {'metadata': {'model': {'input_shapes': shapes, 'output_slices': {}}},
|
||||
'run_model': {(CAM_W, CAM_H): tests_helpers._noop_jit}}
|
||||
with open(tmp_path / 'driving_test_tinygrad.pkl', 'wb') as f:
|
||||
dump_oob(pkl_data, f)
|
||||
bundle = DummyBundle(models=[DummyModel('supercombo', 'driving_test_tinygrad.pkl')])
|
||||
patch_modeld(bundle)
|
||||
monkeypatch.setattr(hw.Paths, 'model_root', staticmethod(lambda: str(tmp_path)))
|
||||
state = ModelState(cam_w=CAM_W, cam_h=CAM_H)
|
||||
assert state.is_run_model and state.run_model is not None
|
||||
assert state.run_policy is None and state.warp is None
|
||||
assert 'img' in state.frame_views and 'big_img' in state.frame_views
|
||||
|
||||
|
||||
ARCHETYPE_NAMES = list(ARCHETYPES.keys())
|
||||
|
||||
|
||||
@@ -1,81 +0,0 @@
|
||||
import numpy as np
|
||||
from openpilot.common.test import OpenpilotTestCase
|
||||
from openpilot.sunnypilot.modeld_v2.constants import ModelConstants
|
||||
from openpilot.sunnypilot.modeld_v2.parse_model_outputs import Parser, _infer_mhp, sigmoid, softmax
|
||||
|
||||
|
||||
class TestParseModelOutputs(OpenpilotTestCase):
|
||||
def test_infer_mhp_lead(self):
|
||||
in_hypotheses, out_selections = _infer_mhp(102, 24)
|
||||
assert in_hypotheses == 2
|
||||
assert out_selections == 3
|
||||
|
||||
def test_infer_mhp_plan(self):
|
||||
in_hypotheses, out_selections = _infer_mhp(4955, 495)
|
||||
assert in_hypotheses == 5
|
||||
assert out_selections == 1
|
||||
|
||||
def test_infer_mhp_non_mdn(self):
|
||||
in_hypotheses, out_selections = _infer_mhp(48, 24)
|
||||
assert in_hypotheses == 1
|
||||
assert out_selections == 0
|
||||
|
||||
def test_check_missing_raises(self):
|
||||
parser = Parser(ignore_missing=False)
|
||||
with self.assertRaises(ValueError):
|
||||
parser.check_missing({}, "missing_key")
|
||||
|
||||
def test_check_missing_ignored(self):
|
||||
parser = Parser(ignore_missing=True)
|
||||
assert parser.check_missing({}, "missing_key") is True
|
||||
|
||||
def test_binary_crossentropy(self):
|
||||
parser = Parser()
|
||||
raw_logits = np.array([[-10.0, 0.0, 10.0]], dtype=np.float32)
|
||||
outs = {"meta": raw_logits.copy()}
|
||||
parser.parse_binary_crossentropy("meta", outs)
|
||||
expected_probabilities = sigmoid(raw_logits)
|
||||
np.testing.assert_allclose(outs["meta"], expected_probabilities, rtol=1e-5, atol=1e-6)
|
||||
|
||||
def test_categorical_crossentropy(self):
|
||||
parser = Parser()
|
||||
raw_logits = np.array([[1.0, 2.0, 3.0]], dtype=np.float32)
|
||||
outs = {"desire_state": raw_logits.copy()}
|
||||
parser.parse_categorical_crossentropy("desire_state", outs)
|
||||
expected_probabilities = softmax(raw_logits)
|
||||
np.testing.assert_allclose(outs["desire_state"], expected_probabilities, rtol=1e-5, atol=1e-6)
|
||||
|
||||
def test_parse_vision_outputs(self):
|
||||
parser = Parser()
|
||||
pose_raw = np.zeros((1, ModelConstants.POSE_WIDTH * 2), dtype=np.float32)
|
||||
road_transform_raw = np.zeros((1, ModelConstants.POSE_WIDTH * 2), dtype=np.float32)
|
||||
lead_raw = np.zeros((1, 102), dtype=np.float32)
|
||||
meta_raw = np.zeros((1, 55), dtype=np.float32)
|
||||
vision_outputs = {"pose": pose_raw, "road_transform": road_transform_raw, "lead": lead_raw, "meta": meta_raw}
|
||||
parsed = parser.parse_vision_outputs(vision_outputs)
|
||||
assert "pose" in parsed
|
||||
assert "road_transform" in parsed
|
||||
assert "lead" in parsed
|
||||
assert "meta" in parsed
|
||||
assert parsed["pose"].shape == (1, ModelConstants.POSE_WIDTH)
|
||||
assert parsed["lead"].shape == (1, ModelConstants.LEAD_MHP_SELECTION, ModelConstants.LEAD_TRAJ_LEN, ModelConstants.LEAD_WIDTH)
|
||||
|
||||
def test_parse_policy_outputs(self):
|
||||
parser = Parser()
|
||||
plan_raw = np.zeros((1, 4955), dtype=np.float32)
|
||||
desire_state_raw = np.zeros((1, ModelConstants.DESIRE_PRED_WIDTH), dtype=np.float32)
|
||||
action_raw = np.zeros((1, ModelConstants.ACTION_WIDTH * 2), dtype=np.float32)
|
||||
policy_outputs = {"plan": plan_raw, "desire_state": desire_state_raw, "action": action_raw}
|
||||
parsed = parser.parse_policy_outputs(policy_outputs)
|
||||
assert parsed["plan"].shape == (1, ModelConstants.IDX_N, ModelConstants.PLAN_WIDTH)
|
||||
assert parsed["action"].shape == (1, ModelConstants.ACTION_WIDTH)
|
||||
assert parsed["desire_state"].shape == (1, ModelConstants.DESIRE_PRED_WIDTH)
|
||||
|
||||
def test_parse_outputs_combined(self):
|
||||
parser = Parser()
|
||||
outputs = {"plan": np.zeros((1, 4955), dtype=np.float32), "pose": np.zeros((1, ModelConstants.POSE_WIDTH * 2),
|
||||
dtype=np.float32), "meta": np.zeros((1, 55), dtype=np.float32)}
|
||||
parsed = parser.parse_outputs(outputs)
|
||||
assert parsed["plan"].shape == (1, ModelConstants.IDX_N, ModelConstants.PLAN_WIDTH)
|
||||
assert parsed["pose"].shape == (1, ModelConstants.POSE_WIDTH)
|
||||
assert parsed["meta"].shape == (1, 55)
|
||||
@@ -0,0 +1,121 @@
|
||||
import numpy as np
|
||||
|
||||
def index_function(idx, max_val=192, max_idx=32):
|
||||
return max_val * ((idx/max_idx)**2)
|
||||
|
||||
|
||||
class ModelConstants:
|
||||
# time and distance indices
|
||||
IDX_N = 33
|
||||
T_IDXS = [index_function(idx, max_val=10.0) for idx in range(IDX_N)]
|
||||
X_IDXS = [index_function(idx, max_val=192.0) for idx in range(IDX_N)]
|
||||
LEAD_T_IDXS = [0., 2., 4., 6., 8., 10.]
|
||||
LEAD_T_OFFSETS = [0., 2., 4.]
|
||||
META_T_IDXS = [2., 4., 6., 8., 10.]
|
||||
|
||||
# model inputs constants
|
||||
MODEL_FREQ = 20
|
||||
FEATURE_LEN = 512
|
||||
HISTORY_BUFFER_LEN = 99
|
||||
DESIRE_LEN = 8
|
||||
TRAFFIC_CONVENTION_LEN = 2
|
||||
NAV_FEATURE_LEN = 256
|
||||
NAV_INSTRUCTION_LEN = 150
|
||||
LAT_PLANNER_STATE_LEN = 4
|
||||
LATERAL_CONTROL_PARAMS_LEN = 2
|
||||
PREV_DESIRED_CURV_LEN = 1
|
||||
|
||||
# model outputs constants
|
||||
FCW_THRESHOLDS_5MS2 = np.array([.05, .05, .15, .15, .15], dtype=np.float32)
|
||||
FCW_THRESHOLDS_3MS2 = np.array([.7, .7], dtype=np.float32)
|
||||
FCW_5MS2_PROBS_WIDTH = 5
|
||||
FCW_3MS2_PROBS_WIDTH = 2
|
||||
|
||||
DISENGAGE_WIDTH = 5
|
||||
POSE_WIDTH = 6
|
||||
WIDE_FROM_DEVICE_WIDTH = 3
|
||||
SIM_POSE_WIDTH = 6
|
||||
LEAD_WIDTH = 4
|
||||
LANE_LINES_WIDTH = 2
|
||||
ROAD_EDGES_WIDTH = 2
|
||||
PLAN_WIDTH = 15
|
||||
DESIRE_PRED_WIDTH = 8
|
||||
LAT_PLANNER_SOLUTION_WIDTH = 4
|
||||
DESIRED_CURV_WIDTH = 1
|
||||
|
||||
NUM_LANE_LINES = 4
|
||||
NUM_ROAD_EDGES = 2
|
||||
|
||||
LEAD_TRAJ_LEN = 6
|
||||
DESIRE_PRED_LEN = 4
|
||||
|
||||
PLAN_MHP_N = 5
|
||||
LEAD_MHP_N = 2
|
||||
PLAN_MHP_SELECTION = 1
|
||||
LEAD_MHP_SELECTION = 3
|
||||
|
||||
FCW_THRESHOLD_5MS2_HIGH = 0.15
|
||||
FCW_THRESHOLD_5MS2_LOW = 0.05
|
||||
FCW_THRESHOLD_3MS2 = 0.7
|
||||
|
||||
CONFIDENCE_BUFFER_LEN = 5
|
||||
RYG_GREEN = 0.01165
|
||||
RYG_YELLOW = 0.06157
|
||||
|
||||
POLY_PATH_DEGREE = 4
|
||||
|
||||
|
||||
# model outputs slices
|
||||
class Plan:
|
||||
POSITION = slice(0, 3)
|
||||
VELOCITY = slice(3, 6)
|
||||
ACCELERATION = slice(6, 9)
|
||||
T_FROM_CURRENT_EULER = slice(9, 12)
|
||||
ORIENTATION_RATE = slice(12, 15)
|
||||
|
||||
|
||||
class Meta:
|
||||
ENGAGED = slice(0, 1)
|
||||
# next 2, 4, 6, 8, 10 seconds
|
||||
GAS_DISENGAGE = slice(1, 31, 6)
|
||||
BRAKE_DISENGAGE = slice(2, 31, 6)
|
||||
STEER_OVERRIDE = slice(3, 31, 6)
|
||||
HARD_BRAKE_3 = slice(4, 31, 6)
|
||||
HARD_BRAKE_4 = slice(5, 31, 6)
|
||||
HARD_BRAKE_5 = slice(6, 31, 6)
|
||||
# next 0, 2, 4, 6, 8, 10 seconds
|
||||
GAS_PRESS = slice(31, 55, 4)
|
||||
BRAKE_PRESS = slice(32, 55, 4)
|
||||
LEFT_BLINKER = slice(33, 55, 4)
|
||||
RIGHT_BLINKER = slice(34, 55, 4)
|
||||
|
||||
|
||||
class MetaTombRaider:
|
||||
ENGAGED = slice(0, 1)
|
||||
# next 2, 4, 6, 8, 10 seconds
|
||||
GAS_DISENGAGE = slice(1, 41, 8)
|
||||
BRAKE_DISENGAGE = slice(2, 41, 8)
|
||||
STEER_OVERRIDE = slice(3, 41, 8)
|
||||
HARD_BRAKE_3 = slice(4, 41, 8)
|
||||
HARD_BRAKE_4 = slice(5, 41, 8)
|
||||
HARD_BRAKE_5 = slice(6, 41, 8)
|
||||
GAS_PRESS = slice(7, 41, 8)
|
||||
BRAKE_PRESS = slice(8, 41, 8)
|
||||
# next 0, 2, 4, 6, 8, 10 seconds
|
||||
LEFT_BLINKER = slice(41, 53, 2)
|
||||
RIGHT_BLINKER = slice(42, 53, 2)
|
||||
|
||||
|
||||
class MetaSimPose:
|
||||
ENGAGED = slice(0, 1)
|
||||
# next 2, 4, 6, 8, 10 seconds
|
||||
GAS_DISENGAGE = slice(1, 36, 7)
|
||||
BRAKE_DISENGAGE = slice(2, 36, 7)
|
||||
STEER_OVERRIDE = slice(3, 36, 7)
|
||||
HARD_BRAKE_3 = slice(4, 36, 7)
|
||||
HARD_BRAKE_4 = slice(5, 36, 7)
|
||||
HARD_BRAKE_5 = slice(6, 36, 7)
|
||||
GAS_PRESS = slice(7, 36, 7)
|
||||
# next 0, 2, 4, 6, 8, 10 seconds
|
||||
LEFT_BLINKER = slice(36, 48, 2)
|
||||
RIGHT_BLINKER = slice(37, 48, 2)
|
||||
@@ -139,7 +139,7 @@ class ModelCache:
|
||||
class ModelFetcher:
|
||||
"""Handles fetching and caching of model data from remote source"""
|
||||
MODEL_URL = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_v22.json"
|
||||
MODEL_URL_CHESTNUT = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_chestnut_v25.json"
|
||||
MODEL_URL_CHESTNUT = "https://raw.githubusercontent.com/sunnypilot/sunnypilot-models/refs/heads/gh-pages/docs/driving_models_chestnut_v23.json"
|
||||
|
||||
MODEL_SOURCES = {
|
||||
"qcom": (MODEL_URL, ""),
|
||||
|
||||
@@ -7,11 +7,14 @@ See the LICENSE.md file in the root directory for more details.
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
import pickle
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
|
||||
from openpilot.cereal import custom
|
||||
from openpilot.common.params import Params
|
||||
from openpilot.common.swaglog import cloudlog
|
||||
from openpilot.sunnypilot.models.constants import Meta, MetaSimPose, MetaTombRaider
|
||||
from openpilot.common.hardware.hw import Paths
|
||||
from openpilot.selfdrive.modeld.helpers import chestnut_present
|
||||
|
||||
@@ -19,6 +22,7 @@ from openpilot.selfdrive.modeld.helpers import chestnut_present
|
||||
REQUIRED_JSON_VERSION = 19
|
||||
|
||||
CUSTOM_MODEL_PATH = Paths.model_root()
|
||||
METADATA_PATH = Path(__file__).parent / '../models/supercombo_metadata.pkl'
|
||||
ModelManager = custom.ModelManagerSP
|
||||
|
||||
ACTIVE_BUNDLE_KEYS = {
|
||||
@@ -197,6 +201,33 @@ def _get_model():
|
||||
return None
|
||||
|
||||
|
||||
def load_metadata():
|
||||
metadata_path = METADATA_PATH
|
||||
|
||||
with open(metadata_path, 'rb') as f:
|
||||
return pickle.load(f)
|
||||
|
||||
|
||||
def prepare_inputs(model_metadata: dict) -> dict[str, np.ndarray]:
|
||||
return {
|
||||
key: np.zeros(shape, dtype=np.float32).flatten()
|
||||
for key, shape in model_metadata['input_shapes'].items()
|
||||
if 'img' not in key
|
||||
}
|
||||
|
||||
|
||||
def load_meta_constants(model_metadata: dict):
|
||||
""" Loads the appropriate meta model class based on key shapes"""
|
||||
if 'sim_pose' in model_metadata['input_shapes']:
|
||||
return MetaSimPose
|
||||
|
||||
meta_slice = model_metadata['output_slices']['meta']
|
||||
if (meta_slice.start, meta_slice.stop, meta_slice.step) == (5868, 5921, None):
|
||||
return MetaTombRaider
|
||||
|
||||
return Meta
|
||||
|
||||
|
||||
# The following method(s) are modeld helper methods
|
||||
def plan_x_idxs_helper(constants, plan, model_output) -> list[float]:
|
||||
# times at X_IDXS according to plan.
|
||||
|
||||
@@ -1,261 +0,0 @@
|
||||
"""
|
||||
Copyright (c) 2021-, Haibin Wen, sunnypilot, and a number of other contributors.
|
||||
|
||||
This file is part of sunnypilot and is licensed under the MIT License.
|
||||
See the LICENSE.md file in the root directory for more details.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import socket
|
||||
import threading
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from openpilot.common.params import Params
|
||||
from openpilot.common.swaglog import cloudlog
|
||||
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_pairing import (
|
||||
BEACON_PREFIX,
|
||||
DISCOVERED_APP_KEY,
|
||||
SUNNYLINK_LOCAL_UDP_PORT,
|
||||
format_endpoint,
|
||||
get_local_apps,
|
||||
pairing_requested,
|
||||
update_local_app_endpoint,
|
||||
)
|
||||
|
||||
LOCAL_BEACON_FRESH_S = 30
|
||||
|
||||
@dataclass
|
||||
class AppBeacon:
|
||||
"""A parsed app beacon — the app announcing it is acting as the local backend."""
|
||||
app_id: str
|
||||
ws_port: int
|
||||
source_ip: str
|
||||
|
||||
@property
|
||||
def endpoint(self) -> str:
|
||||
return format_endpoint(self.source_ip, self.ws_port)
|
||||
|
||||
|
||||
def parse_beacon(raw: str | bytes, source_ip: str = "") -> AppBeacon | None:
|
||||
"""
|
||||
Parse one UDP beacon line from the app.
|
||||
|
||||
Wire format: `SUNNYLINK1 {"v":1,"role":"app","app_id":"<uuid>","ws_port":8443}`
|
||||
Returns None for anything else. Beacons carry ids + addresses only — no secrets.
|
||||
"""
|
||||
if isinstance(raw, bytes):
|
||||
raw = raw.decode("utf-8", errors="replace")
|
||||
raw = raw.strip()
|
||||
if not raw.startswith(BEACON_PREFIX + " "):
|
||||
return None
|
||||
try:
|
||||
data = json.loads(raw[len(BEACON_PREFIX) + 1:])
|
||||
except ValueError:
|
||||
return None
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
if data.get("role") != "app" or data.get("v") != 1:
|
||||
return None
|
||||
app_id = data.get("app_id")
|
||||
ws_port = data.get("ws_port")
|
||||
if not isinstance(app_id, str) or not app_id:
|
||||
return None
|
||||
if not isinstance(ws_port, int) or not (0 < ws_port <= 65535):
|
||||
return None
|
||||
return AppBeacon(app_id=app_id, ws_port=ws_port, source_ip=source_ip)
|
||||
|
||||
|
||||
class LocalDiscovery(threading.Thread):
|
||||
"""
|
||||
Passive UDP listener
|
||||
|
||||
- While a pairing window is armed: track the freshest app beacon so the
|
||||
daemon can offer pairing to a NEW app, and mirror it into a status param
|
||||
for the settings UI.
|
||||
- Independently of any window: a beacon from an app ALREADY in the paired
|
||||
registry refreshes its cached endpoint — IPs are not identity, the app can
|
||||
move between networks.
|
||||
"""
|
||||
|
||||
def __init__(self, params: Params | None = None, port: int = SUNNYLINK_LOCAL_UDP_PORT,
|
||||
sock: socket.socket | None = None, write_interval_s: float = 5.0,
|
||||
paired_refresh_cb: Callable[[AppBeacon], None] | None = None):
|
||||
super().__init__(name="local_discovery_listener", daemon=True)
|
||||
self.params = params or Params()
|
||||
self.port = port
|
||||
self._sock = sock
|
||||
self.paired_refresh_cb = paired_refresh_cb
|
||||
self._latest_endpoint: str | None = None
|
||||
self._latest_app_id: str | None = None
|
||||
self._last_seen_monotonic: float = 0.0
|
||||
self._latest_paired_endpoint: str | None = None
|
||||
self._latest_paired_app_id: str | None = None
|
||||
self._last_paired_seen_monotonic: float = 0.0
|
||||
self._lock = threading.Lock()
|
||||
self._stop_event = threading.Event()
|
||||
self.write_interval_s = write_interval_s
|
||||
self._last_write_monotonic = 0.0
|
||||
self._last_written_endpoint: str | None = None
|
||||
self._last_written_app_id: str | None = None
|
||||
self._discovered_cleared = False
|
||||
|
||||
def stop(self) -> None:
|
||||
self._stop_event.set()
|
||||
if self._sock is not None:
|
||||
try:
|
||||
self._sock.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def latest_endpoint(self) -> str | None:
|
||||
"""The most recently announced app endpoint (None outside a pairing window)."""
|
||||
with self._lock:
|
||||
return self._latest_endpoint
|
||||
|
||||
def latest_app_id(self) -> str | None:
|
||||
"""The app_id of the most recently announced beacon (None outside a window)."""
|
||||
with self._lock:
|
||||
return self._latest_app_id
|
||||
|
||||
def last_seen_ago(self) -> float | None:
|
||||
"""Seconds since the last app beacon was heard (None when none heard yet)."""
|
||||
with self._lock:
|
||||
if self._last_seen_monotonic == 0.0:
|
||||
return None
|
||||
return time.monotonic() - self._last_seen_monotonic
|
||||
|
||||
def latest_paired_endpoint(self) -> str | None:
|
||||
"""The freshest beacon endpoint announced by an ALREADY-PAIRED."""
|
||||
with self._lock:
|
||||
return self._latest_paired_endpoint
|
||||
|
||||
def latest_paired_app_id(self) -> str | None:
|
||||
"""The app_id of the freshest paired-app beacon (None until one is heard)."""
|
||||
with self._lock:
|
||||
return self._latest_paired_app_id
|
||||
|
||||
def latest_paired_seen_ago(self) -> float | None:
|
||||
"""Seconds since the freshest paired-app beacon was heard (None when none)."""
|
||||
with self._lock:
|
||||
if self._last_paired_seen_monotonic == 0.0:
|
||||
return None
|
||||
return time.monotonic() - self._last_paired_seen_monotonic
|
||||
|
||||
def _handle(self, raw: bytes, source_ip: str) -> None:
|
||||
beacon = parse_beacon(raw, source_ip)
|
||||
if beacon is None:
|
||||
return
|
||||
if pairing_requested(self.params):
|
||||
with self._lock:
|
||||
self._latest_endpoint = beacon.endpoint
|
||||
self._latest_app_id = beacon.app_id
|
||||
self._last_seen_monotonic = time.monotonic()
|
||||
self._write_discovered_param(beacon)
|
||||
cloudlog.debug(f"local_discovery.app_found {beacon.app_id} at {beacon.endpoint}")
|
||||
else:
|
||||
with self._lock:
|
||||
self._latest_endpoint = None
|
||||
self._latest_app_id = None
|
||||
self._last_seen_monotonic = 0.0
|
||||
self._clear_discovered_param()
|
||||
self._maybe_refresh_paired_app(beacon)
|
||||
|
||||
def _maybe_refresh_paired_app(self, beacon: AppBeacon) -> None:
|
||||
"""Refresh a paired app's registry endpoint from its beacon."""
|
||||
|
||||
if not any(app.app_id == beacon.app_id for app in get_local_apps(self.params)):
|
||||
return
|
||||
with self._lock:
|
||||
self._latest_paired_endpoint = beacon.endpoint
|
||||
self._latest_paired_app_id = beacon.app_id
|
||||
self._last_paired_seen_monotonic = time.monotonic()
|
||||
if update_local_app_endpoint(beacon.app_id, beacon.endpoint, self.params):
|
||||
cloudlog.debug(f"local_discovery.paired_refresh {beacon.app_id} -> {beacon.endpoint}")
|
||||
if self.paired_refresh_cb is not None:
|
||||
try:
|
||||
self.paired_refresh_cb(beacon)
|
||||
except Exception:
|
||||
cloudlog.exception("local_discovery.paired_refresh_cb.exception")
|
||||
|
||||
def _clear_discovered_param(self) -> None:
|
||||
if self._discovered_cleared:
|
||||
return
|
||||
self._discovered_cleared = True
|
||||
try:
|
||||
self.params.remove(DISCOVERED_APP_KEY)
|
||||
except Exception:
|
||||
cloudlog.exception("local_discovery.param_clear.exception")
|
||||
|
||||
def _write_discovered_param(self, beacon: AppBeacon) -> None:
|
||||
"""Mirror the freshest beacon into a param the settings UI can read."""
|
||||
now = time.monotonic()
|
||||
changed = beacon.endpoint != self._last_written_endpoint or beacon.app_id != self._last_written_app_id
|
||||
if not changed and now - self._last_write_monotonic < self.write_interval_s:
|
||||
return
|
||||
self._last_write_monotonic = now
|
||||
self._last_written_endpoint = beacon.endpoint
|
||||
self._last_written_app_id = beacon.app_id
|
||||
self._discovered_cleared = False
|
||||
payload = {
|
||||
"endpoint": beacon.endpoint,
|
||||
"app_id": beacon.app_id,
|
||||
"ts": int(time.monotonic()),
|
||||
}
|
||||
try:
|
||||
self.params.put(DISCOVERED_APP_KEY, payload, block=True)
|
||||
except Exception:
|
||||
cloudlog.exception("local_discovery.param_write.exception")
|
||||
|
||||
def _bind(self) -> socket.socket:
|
||||
if self._sock is not None:
|
||||
return self._sock
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
sock.bind(("0.0.0.0", self.port))
|
||||
sock.settimeout(0.5)
|
||||
return sock
|
||||
|
||||
def run(self) -> None:
|
||||
sock = self._bind()
|
||||
try:
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
data, addr = sock.recvfrom(4096)
|
||||
self._handle(data, addr[0] if len(addr) > 0 else "")
|
||||
except TimeoutError:
|
||||
continue
|
||||
except OSError:
|
||||
# Socket closed by stop() — exit quietly.
|
||||
if self._stop_event.is_set():
|
||||
break
|
||||
cloudlog.exception("local_discovery.recv.exception")
|
||||
break
|
||||
finally:
|
||||
try:
|
||||
sock.close()
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
|
||||
def latest_discovered_app(params: Params | None = None,
|
||||
fresh_s: float = LOCAL_BEACON_FRESH_S) -> tuple[str, int] | None:
|
||||
params = params or Params()
|
||||
data = params.get(DISCOVERED_APP_KEY)
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
endpoint = str(data.get("endpoint", ""))
|
||||
try:
|
||||
ts = int(data.get("ts") or 0)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
if not endpoint or ts <= 0:
|
||||
return None
|
||||
age = time.monotonic() - ts
|
||||
# Negative age = written before the last reboot (monotonic restarts at boot).
|
||||
if age < 0 or age > fresh_s:
|
||||
return None
|
||||
return endpoint, max(0, int(age))
|
||||
@@ -1,257 +0,0 @@
|
||||
"""
|
||||
Copyright (c) 2021-, Haibin Wen, sunnypilot, and a number of other contributors.
|
||||
|
||||
This file is part of sunnypilot and is licensed under the MIT License.
|
||||
See the LICENSE.md file in the root directory for more details.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
from datetime import datetime, UTC
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from openpilot.common.params import Params
|
||||
from openpilot.common.swaglog import cloudlog
|
||||
|
||||
SUNNYLINK_LOCAL_UDP_PORT = 53133
|
||||
SUNNYLINK_LOCAL_WS_PORT = 8443
|
||||
|
||||
LOCAL_APPS_KEY = "SunnylinkLocalApps"
|
||||
PAIRING_CODE_KEY = "SunnylinkLocalPairingCode"
|
||||
PAIRING_REQUEST_KEY = "SunnylinkLocalPairingRequest"
|
||||
DISCOVERED_APP_KEY = "SunnylinkLocalDiscoveredApp"
|
||||
|
||||
PAIRING_CODE_LENGTH = 6
|
||||
PAIRING_CODE_ALPHABET = "0123456789"
|
||||
DEFAULT_CODE_ROTATION_S = 10 * 60 # re-roll the displayed code every 10 min
|
||||
PAIRING_WINDOW_S = 5 * 60
|
||||
|
||||
BEACON_PREFIX = "SUNNYLINK1"
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalApp:
|
||||
"""One app paired with this device (the app runs the local "backend")."""
|
||||
app_id: str
|
||||
endpoint: str
|
||||
app_name: str = ""
|
||||
alias: str = ""
|
||||
paired_at: int = 0 # epoch seconds
|
||||
|
||||
@staticmethod
|
||||
def from_dict(data: dict[str, Any]) -> LocalApp:
|
||||
return LocalApp(
|
||||
app_id=str(data.get("app_id", "")),
|
||||
endpoint=str(data.get("endpoint", "")),
|
||||
app_name=str(data.get("app_name", "")),
|
||||
alias=str(data.get("alias", "")),
|
||||
paired_at=int(data.get("paired_at") or 0),
|
||||
)
|
||||
|
||||
|
||||
def local_app_display_name(app: LocalApp) -> str:
|
||||
return app.alias or app.app_name or app.app_id
|
||||
|
||||
|
||||
def is_locally_paired(params: Params | None = None) -> bool:
|
||||
return len(get_local_apps(params)) > 0
|
||||
|
||||
|
||||
def get_local_apps(params: Params | None = None) -> list[LocalApp]:
|
||||
"""The paired-app registry (a JSON list persisted in `SunnylinkLocalApps`)."""
|
||||
params = params or Params()
|
||||
data = params.get(LOCAL_APPS_KEY)
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return [LocalApp.from_dict(item) for item in data if isinstance(item, dict) and item.get("app_id")]
|
||||
|
||||
|
||||
def _save_local_apps(apps: list[LocalApp], params: Params | None = None) -> None:
|
||||
params = params or Params()
|
||||
if apps:
|
||||
params.put(LOCAL_APPS_KEY, [asdict(app) for app in apps], block=True)
|
||||
else:
|
||||
params.remove(LOCAL_APPS_KEY)
|
||||
|
||||
|
||||
def add_local_app(app: LocalApp, params: Params | None = None) -> None:
|
||||
"""Append (or update by app_id) and persist."""
|
||||
if not app.paired_at:
|
||||
app.paired_at = int(datetime.now(UTC).replace(tzinfo=None).timestamp())
|
||||
apps = [existing for existing in get_local_apps(params) if existing.app_id != app.app_id]
|
||||
apps.append(app)
|
||||
_save_local_apps(apps, params)
|
||||
cloudlog.event("local_pairing.app_paired", app_id=app.app_id, endpoint=app.endpoint)
|
||||
|
||||
|
||||
def update_local_app_endpoint(app_id: str, endpoint: str, params: Params | None = None) -> bool:
|
||||
"""Refresh a PAIRED app's cached endpoint from its beacon."""
|
||||
apps = get_local_apps(params)
|
||||
for i, app in enumerate(apps):
|
||||
if app.app_id != app_id or app.endpoint == endpoint:
|
||||
continue
|
||||
apps[i] = LocalApp(app_id=app.app_id, endpoint=endpoint,
|
||||
app_name=app.app_name, alias=app.alias, paired_at=app.paired_at)
|
||||
_save_local_apps(apps, params)
|
||||
cloudlog.event("local_pairing.app_endpoint_refreshed", app_id=app_id, endpoint=endpoint)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def set_local_app_alias(app_id: str, alias: str, params: Params | None = None) -> bool:
|
||||
apps = get_local_apps(params)
|
||||
for i, app in enumerate(apps):
|
||||
if app.app_id != app_id:
|
||||
continue
|
||||
if app.alias == alias:
|
||||
return False
|
||||
apps[i] = LocalApp(app_id=app.app_id, endpoint=app.endpoint,
|
||||
app_name=app.app_name, alias=alias, paired_at=app.paired_at)
|
||||
_save_local_apps(apps, params)
|
||||
cloudlog.event("local_pairing.app_alias_updated", app_id=app_id, alias=alias)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def remove_local_app(app_id: str, params: Params | None = None) -> bool:
|
||||
"""Unpair an app by id. Returns True when an app was removed."""
|
||||
apps = get_local_apps(params)
|
||||
remaining = [app for app in apps if app.app_id != app_id]
|
||||
if len(remaining) == len(apps):
|
||||
return False
|
||||
_save_local_apps(remaining, params)
|
||||
cloudlog.event("local_pairing.app_unpaired", app_id=app_id)
|
||||
return True
|
||||
|
||||
|
||||
def remove_all_local_apps(params: Params | None = None) -> None:
|
||||
"""Unpair every app."""
|
||||
_save_local_apps([], params)
|
||||
cloudlog.event("local_pairing.all_apps_unpaired")
|
||||
|
||||
|
||||
def generate_pairing_code() -> str:
|
||||
"""A 6-digit numeric pairing code."""
|
||||
return "".join(secrets.choice(PAIRING_CODE_ALPHABET) for _ in range(PAIRING_CODE_LENGTH))
|
||||
|
||||
|
||||
def _write_pairing_code(code: str, params: Params) -> None:
|
||||
"""Persist the code with its armed-at monotonic timestamp — the window is derived from it."""
|
||||
params.put(PAIRING_CODE_KEY, {"code": code, "ts": int(time.monotonic())}, block=True)
|
||||
|
||||
|
||||
def read_pairing_code(params: Params | None = None) -> str | None:
|
||||
"""The stored pairing code, or None when cleared / not yet generated."""
|
||||
params = params or Params()
|
||||
data = params.get(PAIRING_CODE_KEY)
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
code = data.get("code")
|
||||
return str(code) if code else None
|
||||
|
||||
|
||||
def get_pairing_code(params: Params | None = None) -> str:
|
||||
"""The displayed pairing code, generating and persisting one on first use."""
|
||||
params = params or Params()
|
||||
code = read_pairing_code(params)
|
||||
if code is None:
|
||||
code = generate_pairing_code()
|
||||
_write_pairing_code(code, params)
|
||||
return code
|
||||
|
||||
|
||||
def pairing_requested(params: Params | None = None) -> bool:
|
||||
"""True while the pairing window is armed and fresh.
|
||||
|
||||
Self-expiring: if the code (which carries the armed-at timestamp) is missing
|
||||
or older than PAIRING_WINDOW_S, the flag is dropped here."""
|
||||
params = params or Params()
|
||||
if not params.get_bool(PAIRING_REQUEST_KEY):
|
||||
return False
|
||||
data = params.get(PAIRING_CODE_KEY)
|
||||
ts = data.get("ts") if isinstance(data, dict) else None
|
||||
if not isinstance(ts, (int, float)):
|
||||
clear_pairing_request(params)
|
||||
return False
|
||||
age = time.monotonic() - ts
|
||||
# Negative age = armed before the last reboot (monotonic restarts at boot).
|
||||
if age < 0 or age > PAIRING_WINDOW_S:
|
||||
clear_pairing_request(params)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def arm_pairing(params: Params | None = None) -> str:
|
||||
"""Arm a pairing window and return the code for the app.
|
||||
|
||||
Rolls a fresh code first, then sets the flag, so pairing_requested never
|
||||
sees an armed flag without a valid code."""
|
||||
params = params or Params()
|
||||
code = generate_pairing_code()
|
||||
_write_pairing_code(code, params)
|
||||
params.put_bool(PAIRING_REQUEST_KEY, True, block=True)
|
||||
return code
|
||||
|
||||
|
||||
def clear_pairing_request(params: Params | None = None) -> None:
|
||||
"""Close the pairing window: drop the request flag and the code together."""
|
||||
params = params or Params()
|
||||
params.remove(PAIRING_REQUEST_KEY)
|
||||
params.remove(PAIRING_CODE_KEY)
|
||||
|
||||
|
||||
def verify_pairing_code(code: str, params: Params | None = None) -> bool:
|
||||
"""Constant-time check of a code typed into the app against the displayed one."""
|
||||
params = params or Params()
|
||||
current = read_pairing_code(params)
|
||||
if current is None:
|
||||
return False
|
||||
return secrets.compare_digest(str(code).strip().upper(), current)
|
||||
|
||||
|
||||
class PairingCodeRotator(threading.Thread):
|
||||
"""Re-roll the displayed code while a pairing window is armed; clear it
|
||||
otherwise — the code is never generated outside a window."""
|
||||
|
||||
def __init__(self, params: Params | None = None, rotation_s: float = DEFAULT_CODE_ROTATION_S,
|
||||
stop_event: threading.Event | None = None, tick_cb: Callable[[], None] | None = None):
|
||||
super().__init__(name="local_pairing_code_rotator", daemon=True)
|
||||
self.params = params or Params()
|
||||
self.rotation_s = rotation_s
|
||||
self.stop_event = stop_event or threading.Event()
|
||||
# Test seam: invoked once per loop iteration after state is updated.
|
||||
self.tick_cb = tick_cb
|
||||
|
||||
def rotate(self) -> None:
|
||||
"""Re-roll the code while the window is armed, clear it otherwise."""
|
||||
if pairing_requested(self.params):
|
||||
_write_pairing_code(generate_pairing_code(), self.params)
|
||||
else:
|
||||
self.params.remove(PAIRING_CODE_KEY)
|
||||
|
||||
def run(self) -> None:
|
||||
self.rotate()
|
||||
while not self.stop_event.wait(self.rotation_s):
|
||||
try:
|
||||
self.rotate()
|
||||
if self.tick_cb is not None:
|
||||
self.tick_cb()
|
||||
except Exception:
|
||||
cloudlog.exception("local_pairing.code_rotator.exception")
|
||||
|
||||
|
||||
def format_endpoint(host: str, ws_port: int = SUNNYLINK_LOCAL_WS_PORT) -> str:
|
||||
return f"ws://{host}:{ws_port}"
|
||||
|
||||
|
||||
def local_identity(params: Params | None = None) -> str:
|
||||
"""Identity claim on local connections. DongleId always exists on comma
|
||||
hardware (SunnylinkDongleId is "UnregisteredDevice" until cloud
|
||||
registration) and is what the app matches against the backend device list
|
||||
to dedupe cloud + local entries."""
|
||||
params = params or Params()
|
||||
return params.get("DongleId") or params.get("HardwareSerial") or ""
|
||||
@@ -30,26 +30,10 @@ from websocket import (ABNF, WebSocket, WebSocketException, WebSocketTimeoutExce
|
||||
import openpilot.cereal.messaging as messaging
|
||||
from openpilot.sunnypilot.models.model_name import DEFAULT_MODEL, DEFAULT_BIG_MODEL
|
||||
from openpilot.sunnypilot.selfdrive.car.sync_sunnylink_params import update_car_list_param
|
||||
from openpilot.system.athena import rpc as rpc_module
|
||||
from openpilot.sunnypilot.sunnylink.api import SunnylinkApi
|
||||
from openpilot.sunnypilot.sunnylink.utils import sunnylink_need_register, sunnylink_ready, get_param_as_byte, save_param_from_base64_encoded_string
|
||||
from openpilot.sunnypilot.sunnylink.capabilities import generate_capabilities, CAPABILITY_LABELS
|
||||
from openpilot.sunnypilot.sunnylink.tools.generate_settings_schema import generate_schema
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_discovery import LOCAL_BEACON_FRESH_S, AppBeacon, LocalDiscovery
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_pairing import (
|
||||
PAIRING_WINDOW_S,
|
||||
LocalApp,
|
||||
PairingCodeRotator,
|
||||
add_local_app,
|
||||
clear_pairing_request,
|
||||
get_local_apps,
|
||||
is_locally_paired,
|
||||
local_identity,
|
||||
pairing_requested,
|
||||
remove_local_app,
|
||||
set_local_app_alias,
|
||||
verify_pairing_code,
|
||||
)
|
||||
|
||||
SUNNYLINK_ATHENA_HOST = os.getenv('SUNNYLINK_ATHENA_HOST', 'wss://athena.sunnylink.ai')
|
||||
HANDLER_THREADS = int(os.getenv('HANDLER_THREADS', "4"))
|
||||
@@ -58,15 +42,6 @@ SUNNYLINK_LOG_ATTR_NAME = "user.sunny.upload"
|
||||
SUNNYLINK_RECONNECT_TIMEOUT_S = 70 # FYI changing this will also would require a change on sidebar.cc
|
||||
DISALLOW_LOG_UPLOAD = threading.Event()
|
||||
|
||||
LOCAL_PAIRING_SESSION_TIMEOUT_S = PAIRING_WINDOW_S
|
||||
LOCAL_PROBE_INTERVAL_S = 60
|
||||
LOCAL_ENDPOINT_BACKOFF_S = 300
|
||||
PAIRING_WATCHDOG_INTERVAL_S = 2.0
|
||||
|
||||
_active_local_endpoint: str | None = None
|
||||
_active_ws: WebSocket | None = None
|
||||
_pairing_in_progress = threading.Event()
|
||||
|
||||
params = Params()
|
||||
|
||||
# Parameters that should never be remotely modified
|
||||
@@ -291,302 +266,44 @@ def startLocalProxy(global_end_event: threading.Event, remote_ws_uri: str, local
|
||||
return start_local_proxy_shim(global_end_event, local_port, ws)
|
||||
|
||||
|
||||
@dispatcher.add_method
|
||||
def pairLocalApp(code: str, app_id: str = "", app_name: str = "", alias: str = "") -> dict[str, bool | str]:
|
||||
"""Complete pairing with the app on the CURRENT local connection."""
|
||||
if _active_local_endpoint is None:
|
||||
return {"success": False, "error": "not connected to a local app"}
|
||||
if not verify_pairing_code(code):
|
||||
cloudlog.warning("sunnylinkd.pairLocalApp.invalid_code")
|
||||
return {"success": False, "error": "invalid code"}
|
||||
add_local_app(LocalApp(app_id=app_id or f"app@{_active_local_endpoint}",
|
||||
endpoint=_active_local_endpoint, app_name=app_name, alias=alias))
|
||||
clear_pairing_request()
|
||||
return {"success": True}
|
||||
|
||||
|
||||
@dispatcher.add_method
|
||||
def updateLocalAppAlias(app_id: str, alias: str) -> dict[str, bool | str]:
|
||||
if _active_local_endpoint is None:
|
||||
return {"success": False, "error": "not connected to a local app"}
|
||||
updated = set_local_app_alias(app_id, alias)
|
||||
return {"success": True, "updated": updated}
|
||||
|
||||
|
||||
@dispatcher.add_method
|
||||
def unpairLocalApp(app_id: str) -> dict[str, bool | str]:
|
||||
if _active_local_endpoint is None:
|
||||
return {"success": False, "error": "not connected to a local app"}
|
||||
removed = remove_local_app(app_id)
|
||||
return {"success": True, "removed": removed}
|
||||
|
||||
|
||||
def _auth_header(is_local: bool) -> dict[str, str]:
|
||||
"""Bearer header for a dial."""
|
||||
api = SunnylinkApi(params.get("SunnylinkDongleId"))
|
||||
payload = {"identity": local_identity()} if is_local else None
|
||||
return {"Authorization": f"Bearer {api.get_token(payload_extra=payload)}"}
|
||||
|
||||
|
||||
def _pairing_session(ws: WebSocket, timeout_s: float = LOCAL_PAIRING_SESSION_TIMEOUT_S) -> bool:
|
||||
"""Serve only the pairing RPCs to an app that isn't in the registry yet —
|
||||
everything else is refused. Returns True if pairing completed (the connection may then
|
||||
serve normally)."""
|
||||
cloudlog.info("sunnylinkd.pairing_session.started")
|
||||
ws.settimeout(10)
|
||||
deadline = time.monotonic() + timeout_s
|
||||
try:
|
||||
while time.monotonic() < deadline and pairing_requested():
|
||||
try:
|
||||
raw = ws.recv() # auto-pongs pings; blocks up to the socket timeout
|
||||
except WebSocketTimeoutException:
|
||||
continue
|
||||
except Exception as e:
|
||||
cloudlog.warning(f"sunnylinkd.pairing_session.{type(e).__name__}")
|
||||
return is_locally_paired()
|
||||
try:
|
||||
msg = rpc_module.loads(raw)
|
||||
except Exception:
|
||||
continue
|
||||
if not rpc_module.is_call(msg):
|
||||
continue
|
||||
if msg.get("method") not in ("pairLocalApp", "unpairLocalApp"):
|
||||
continue # refuse anything but pairing until paired
|
||||
try:
|
||||
ws.send(rpc_module.handle(msg, dispatcher))
|
||||
except Exception as e:
|
||||
cloudlog.warning(f"sunnylinkd.pairing_session.{type(e).__name__}")
|
||||
return is_locally_paired()
|
||||
return is_locally_paired()
|
||||
finally:
|
||||
ws.settimeout(SUNNYLINK_RECONNECT_TIMEOUT_S)
|
||||
|
||||
|
||||
def _pick_ws_uri(discovery: LocalDiscovery, backoffs: dict[str, float]) -> tuple[str, str]:
|
||||
now = time.monotonic()
|
||||
apps = get_local_apps()
|
||||
app_ids = {app.app_id for app in apps}
|
||||
if pairing_requested():
|
||||
latest = discovery.latest_endpoint()
|
||||
seen = discovery.last_seen_ago()
|
||||
app_id = discovery.latest_app_id()
|
||||
if latest is not None and seen is not None and seen <= LOCAL_BEACON_FRESH_S \
|
||||
and app_id is not None and app_id not in app_ids \
|
||||
and backoffs.get(latest, 0.0) <= now:
|
||||
return latest, "pairing_offer"
|
||||
return SUNNYLINK_ATHENA_HOST, "cloud"
|
||||
fresh_endpoint = discovery.latest_paired_endpoint()
|
||||
fresh_seen = discovery.latest_paired_seen_ago()
|
||||
fresh_app_id = discovery.latest_paired_app_id()
|
||||
fresh_ok = (fresh_endpoint is not None and fresh_seen is not None
|
||||
and fresh_seen <= LOCAL_BEACON_FRESH_S and fresh_app_id is not None
|
||||
and fresh_app_id in app_ids)
|
||||
if fresh_ok and backoffs.get(fresh_endpoint, 0.0) <= now:
|
||||
return fresh_endpoint, "paired_local"
|
||||
for app in reversed(apps):
|
||||
if fresh_ok and app.app_id == fresh_app_id:
|
||||
continue
|
||||
if backoffs.get(app.endpoint, 0.0) <= now:
|
||||
return app.endpoint, "paired_local"
|
||||
return SUNNYLINK_ATHENA_HOST, "cloud"
|
||||
|
||||
|
||||
def _probe_local_apps(active_ws: WebSocket, discovery: LocalDiscovery,
|
||||
backoffs: dict[str, float], stop_event: threading.Event) -> None:
|
||||
while not stop_event.wait(LOCAL_PROBE_INTERVAL_S):
|
||||
if pairing_requested():
|
||||
# The pairing watchdog owns an armed window; don't migrate mid-window.
|
||||
continue
|
||||
now = time.monotonic()
|
||||
candidate: str | None = None
|
||||
for app in reversed(get_local_apps()):
|
||||
if backoffs.get(app.endpoint, 0.0) <= now:
|
||||
candidate = app.endpoint
|
||||
break
|
||||
if candidate is None:
|
||||
continue
|
||||
try:
|
||||
probe = create_connection(candidate, header=_auth_header(is_local=True), timeout=10)
|
||||
probe.close()
|
||||
except Exception:
|
||||
backoffs[candidate] = now + LOCAL_ENDPOINT_BACKOFF_S
|
||||
continue
|
||||
cloudlog.event("sunnylinkd.local_probe.reachable", endpoint=candidate)
|
||||
try:
|
||||
active_ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
break
|
||||
|
||||
|
||||
def _handle_paired_refresh(backoffs: dict[str, float], force_attempts: dict[str, float],
|
||||
beacon: AppBeacon) -> None:
|
||||
if not any(app.app_id == beacon.app_id for app in get_local_apps()):
|
||||
return
|
||||
for app in get_local_apps():
|
||||
if app.app_id == beacon.app_id:
|
||||
backoffs.pop(app.endpoint, None)
|
||||
backoffs.pop(beacon.endpoint, None)
|
||||
if _pairing_in_progress.is_set():
|
||||
return
|
||||
if _active_local_endpoint == beacon.endpoint:
|
||||
return
|
||||
if _active_local_endpoint is not None:
|
||||
return
|
||||
now = time.monotonic()
|
||||
if force_attempts.get(beacon.endpoint, 0.0) + LOCAL_BEACON_FRESH_S > now:
|
||||
return
|
||||
force_attempts[beacon.endpoint] = now
|
||||
ws = _active_ws
|
||||
if ws is not None:
|
||||
cloudlog.event("sunnylinkd.paired_refresh.reconnect",
|
||||
app_id=beacon.app_id, endpoint=beacon.endpoint)
|
||||
try:
|
||||
ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _pairing_watchdog(active_ws: WebSocket, backoffs: dict[str, float],
|
||||
stop_event: threading.Event,
|
||||
interval_s: float = PAIRING_WATCHDOG_INTERVAL_S) -> None:
|
||||
"""Watch for the pairing window being armed mid-session and force a
|
||||
re-selection to the newly-discovered app."""
|
||||
while not stop_event.wait(interval_s):
|
||||
if not pairing_requested():
|
||||
continue
|
||||
cloudlog.event("sunnylinkd.pairing_watchdog.arm_detected")
|
||||
for key in list(backoffs):
|
||||
backoffs.pop(key, None)
|
||||
try:
|
||||
active_ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
break
|
||||
|
||||
|
||||
def main(exit_event: threading.Event | None = None):
|
||||
try:
|
||||
set_core_affinity([0, 1, 2, 3])
|
||||
except Exception:
|
||||
cloudlog.exception("failed to set core affinity")
|
||||
|
||||
discovery = LocalDiscovery()
|
||||
code_rotator = PairingCodeRotator()
|
||||
discovery.start()
|
||||
code_rotator.start()
|
||||
|
||||
try:
|
||||
_connection_loop(exit_event, discovery)
|
||||
finally:
|
||||
discovery.stop()
|
||||
code_rotator.stop_event.set()
|
||||
|
||||
|
||||
def _serviceable(params: Params) -> bool:
|
||||
"""sunnylinkd should run when sunnylink is enabled and not on a temporary
|
||||
fault. This deliberately includes the unregistered/unpaired state so a
|
||||
never-registered device can still be discovered and paired over the LAN (the
|
||||
actual session gates — registration/local pairing — are handled per
|
||||
connection inside the loop)."""
|
||||
return params.get_bool("SunnylinkEnabled") and not params.get_bool("SunnylinkTempFault")
|
||||
|
||||
|
||||
def _connection_loop(exit_event: threading.Event | None, discovery: LocalDiscovery) -> None:
|
||||
"""Local-first, cloud-fallback connection loop: a paired local endpoint
|
||||
first, cloud when unreachable, and a pairing session to a freshly-discovered
|
||||
app when a window is armed."""
|
||||
global _active_local_endpoint, _active_ws
|
||||
while sunnylink_need_register(params):
|
||||
cloudlog.info("Waiting for sunnylink registration to complete")
|
||||
time.sleep(10)
|
||||
|
||||
sunnylink_dongle_id = params.get("SunnylinkDongleId")
|
||||
sunnylink_api = SunnylinkApi(sunnylink_dongle_id)
|
||||
UploadQueueCache.initialize(upload_queue)
|
||||
|
||||
update_car_list_param()
|
||||
|
||||
ws_uri = f"{SUNNYLINK_ATHENA_HOST}"
|
||||
conn_start = None
|
||||
conn_retries = 0
|
||||
backoffs: dict[str, float] = {}
|
||||
force_attempts: dict[str, float] = {}
|
||||
discovery.paired_refresh_cb = partial(_handle_paired_refresh, backoffs, force_attempts)
|
||||
|
||||
while (exit_event is None or not exit_event.is_set()) and _serviceable(params):
|
||||
ws_uri, kind = _pick_ws_uri(discovery, backoffs)
|
||||
|
||||
if kind == "cloud" and pairing_requested():
|
||||
cloudlog.debug("sunnylinkd.main.pairing_waiting_for_beacon")
|
||||
time.sleep(3)
|
||||
continue
|
||||
|
||||
if kind == "cloud" and sunnylink_need_register(params):
|
||||
cloudlog.info("Waiting for sunnylink registration or local pairing to complete")
|
||||
time.sleep(10)
|
||||
continue
|
||||
|
||||
if conn_start is None:
|
||||
conn_start = time.monotonic()
|
||||
|
||||
cloudlog.event("sunnylinkd.main.connecting_ws", ws_uri=ws_uri, kind=kind, retries=conn_retries)
|
||||
while (exit_event is None or not exit_event.is_set()) and sunnylink_ready(params):
|
||||
try:
|
||||
if conn_start is None:
|
||||
conn_start = time.monotonic()
|
||||
|
||||
cloudlog.event("sunnylinkd.main.connecting_ws", ws_uri=ws_uri, retries=conn_retries)
|
||||
ws = create_connection(
|
||||
ws_uri,
|
||||
header=_auth_header(is_local=kind != "cloud"),
|
||||
header={"Authorization": f"Bearer {sunnylink_api.get_token()}"},
|
||||
enable_multithread=True,
|
||||
sslopt={"cert_reqs": ssl.CERT_NONE if "localhost" in ws_uri else ssl.CERT_REQUIRED},
|
||||
timeout=SUNNYLINK_RECONNECT_TIMEOUT_S,
|
||||
)
|
||||
except Exception as e:
|
||||
if kind != "cloud":
|
||||
backoffs[ws_uri] = time.monotonic() + LOCAL_ENDPOINT_BACKOFF_S
|
||||
conn_retries += 1
|
||||
params.remove("LastSunnylinkPingTime")
|
||||
_log_connection_error(e)
|
||||
time.sleep(backoff(conn_retries))
|
||||
continue
|
||||
cloudlog.event("sunnylinkd.main.connected_ws", ws_uri=ws_uri, retries=conn_retries,
|
||||
duration=time.monotonic() - conn_start)
|
||||
conn_start = None
|
||||
|
||||
cloudlog.event("sunnylinkd.main.connected_ws", ws_uri=ws_uri, kind=kind, retries=conn_retries,
|
||||
duration=time.monotonic() - conn_start)
|
||||
conn_start = None
|
||||
conn_retries = 0
|
||||
cur_upload_items.clear()
|
||||
_active_ws = ws
|
||||
|
||||
probe_stop: threading.Event | None = None
|
||||
watch_stop: threading.Event | None = None
|
||||
session_endpoint: str | None = ws_uri if kind != "cloud" else None
|
||||
try:
|
||||
if kind == "pairing_offer":
|
||||
_active_local_endpoint = ws_uri
|
||||
_pairing_in_progress.set()
|
||||
try:
|
||||
paired_ok = _pairing_session(ws)
|
||||
finally:
|
||||
_pairing_in_progress.clear()
|
||||
if not paired_ok:
|
||||
backoffs[ws_uri] = time.monotonic() + LOCAL_ENDPOINT_BACKOFF_S
|
||||
conn_retries += 1
|
||||
params.remove("LastSunnylinkPingTime")
|
||||
try:
|
||||
ws.close()
|
||||
except Exception:
|
||||
pass
|
||||
time.sleep(backoff(conn_retries))
|
||||
continue
|
||||
# Paired during the session — this connection may now serve normally.
|
||||
kind = "paired_local"
|
||||
|
||||
if kind == "paired_local":
|
||||
_active_local_endpoint = ws_uri
|
||||
else:
|
||||
_active_local_endpoint = None
|
||||
# While on the cloud link, watch for the local app and migrate back.
|
||||
probe_stop = threading.Event()
|
||||
threading.Thread(target=_probe_local_apps,
|
||||
args=(ws, discovery, backoffs, probe_stop),
|
||||
name="sunnylinkd_local_probe", daemon=True).start()
|
||||
|
||||
# Started after any pairing session on this connection, so it can never
|
||||
# close the connection the code is typed over.
|
||||
watch_stop = threading.Event()
|
||||
threading.Thread(target=_pairing_watchdog, args=(ws, backoffs, watch_stop),
|
||||
name="sunnylinkd_pairing_watchdog", daemon=True).start()
|
||||
conn_retries = 0
|
||||
cur_upload_items.clear()
|
||||
|
||||
handle_long_poll(ws, exit_event)
|
||||
except (KeyboardInterrupt, SystemExit):
|
||||
@@ -594,37 +311,23 @@ def _connection_loop(exit_event: threading.Event | None, discovery: LocalDiscove
|
||||
except Exception as e:
|
||||
conn_retries += 1
|
||||
params.remove("LastSunnylinkPingTime")
|
||||
_log_connection_error(e)
|
||||
finally:
|
||||
if probe_stop is not None:
|
||||
probe_stop.set()
|
||||
if watch_stop is not None:
|
||||
watch_stop.set()
|
||||
if session_endpoint is not None and kind == "paired_local":
|
||||
backoffs.pop(session_endpoint, None)
|
||||
if _active_local_endpoint == session_endpoint:
|
||||
_active_local_endpoint = None
|
||||
if _active_ws is ws:
|
||||
_active_ws = None
|
||||
|
||||
if isinstance(e, (ConnectionError, TimeoutError, WebSocketException)):
|
||||
cloudlog.warning(f"sunnylinkd.main.{type(e).__name__}")
|
||||
elif isinstance(e, OSError):
|
||||
name = errno.errorcode.get(e.errno or -1, "UNKNOWN")
|
||||
msg = f"sunnylinkd.main.OSError.{name} ({e.errno})"
|
||||
is_expected_error = e.errno in (errno.ENETDOWN, errno.ENETRESET, errno.ENETUNREACH)
|
||||
cloudlog.warning(msg) if is_expected_error else cloudlog.exception(msg)
|
||||
else:
|
||||
cloudlog.exception("sunnylinkd.main.exception")
|
||||
|
||||
time.sleep(backoff(conn_retries))
|
||||
|
||||
if not _serviceable(params):
|
||||
cloudlog.debug("Reached end of sunnylinkd.main while sunnylink is not serviceable. Waiting 60s before retrying")
|
||||
if not sunnylink_ready(params):
|
||||
cloudlog.debug("Reached end of sunnylinkd.main while sunnylink is not ready. Waiting 60s before retrying")
|
||||
time.sleep(60)
|
||||
|
||||
|
||||
def _log_connection_error(e: Exception) -> None:
|
||||
if isinstance(e, (ConnectionError, TimeoutError, WebSocketException)):
|
||||
cloudlog.warning(f"sunnylinkd.main.{type(e).__name__}")
|
||||
elif isinstance(e, OSError):
|
||||
name = errno.errorcode.get(e.errno or -1, "UNKNOWN")
|
||||
msg = f"sunnylinkd.main.OSError.{name} ({e.errno})"
|
||||
is_expected_error = e.errno in (errno.ENETDOWN, errno.ENETRESET, errno.ENETUNREACH)
|
||||
cloudlog.warning(msg) if is_expected_error else cloudlog.exception(msg)
|
||||
else:
|
||||
cloudlog.exception("sunnylinkd.main.exception")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
@@ -2,7 +2,6 @@ import base64
|
||||
import gzip
|
||||
import json
|
||||
from openpilot.sunnypilot.sunnylink.api import SunnylinkApi, UNREGISTERED_SUNNYLINK_DONGLE_ID
|
||||
from openpilot.sunnypilot.sunnylink.athena.local_pairing import is_locally_paired
|
||||
from openpilot.common.params import Params, ParamKeyType
|
||||
from openpilot.common.version import is_prebuilt
|
||||
|
||||
@@ -17,11 +16,10 @@ def get_sunnylink_status(params=None) -> tuple[bool, bool, bool]:
|
||||
|
||||
|
||||
def sunnylink_ready(params=None) -> bool:
|
||||
"""Enabled and (cloud-registered or locally paired), and not on a temporary
|
||||
fault. Local pairing makes never-registered devices usable over the LAN."""
|
||||
"""Check if the device is ready to communicate with Sunnylink. That means it is enabled and registered."""
|
||||
params = params or Params()
|
||||
is_sunnylink_enabled, is_registered, is_on_temporary_fault = get_sunnylink_status(params)
|
||||
return is_sunnylink_enabled and (is_registered or is_locally_paired(params)) and not is_on_temporary_fault
|
||||
return is_sunnylink_enabled and is_registered and not is_on_temporary_fault
|
||||
|
||||
|
||||
def use_sunnylink_uploader(params) -> bool:
|
||||
@@ -30,11 +28,10 @@ def use_sunnylink_uploader(params) -> bool:
|
||||
|
||||
|
||||
def sunnylink_need_register(params=None) -> bool:
|
||||
"""Enabled, unregistered, and not locally paired — a locally paired device
|
||||
works without cloud registration and must not be blocked."""
|
||||
"""Check if the device needs to be registered with Sunnylink."""
|
||||
params = params or Params()
|
||||
is_sunnylink_enabled, is_registered, is_on_temporary_fault = get_sunnylink_status(params)
|
||||
return is_sunnylink_enabled and not is_registered and not is_locally_paired(params) and not is_on_temporary_fault
|
||||
return is_sunnylink_enabled and not is_registered and not is_on_temporary_fault
|
||||
|
||||
|
||||
def register_sunnylink():
|
||||
|
||||
+1
-1
Submodule panda updated: 74a0adced4...42643dec15
@@ -38,8 +38,7 @@ def main():
|
||||
api = HfApi()
|
||||
onnx_sha256 = hash_file(args.onnx_path)
|
||||
short_ref = args.onnx_ref[:8]
|
||||
safe_name = args.model_name.replace(" ", "-")
|
||||
folder_name = f"model-{safe_name}-{short_ref}-{args.run_number}"
|
||||
folder_name = f"model-{args.model_name}-{short_ref}-{args.run_number}"
|
||||
|
||||
print(f"ONNX hash: {onnx_sha256}")
|
||||
print(f"ONNX ref: {args.onnx_ref} (short: {short_ref})")
|
||||
|
||||
Reference in New Issue
Block a user