diff --git a/.github/workflows/build-default-big-model.yaml b/.github/workflows/build-default-big-model.yaml new file mode 100644 index 0000000000..4b05a7977d --- /dev/null +++ b/.github/workflows/build-default-big-model.yaml @@ -0,0 +1,83 @@ +name: Build default big model + +on: + workflow_dispatch: + +env: + HF_REPO: sunnypilot/sunnypilot_models_v1 + HF_DEFAULTS_PATH: models/defaults/big + +jobs: + resolve_name: + runs-on: ubuntu-24.04 + outputs: + model_name: ${{ steps.name.outputs.model_name }} + onnx_ref: ${{ steps.name.outputs.onnx_ref }} + steps: + - uses: actions/checkout@v4 + - id: name + run: | + NAME=$(PYTHONPATH=${{ github.workspace }} python3 -c "from openpilot.sunnypilot.models.model_name import DEFAULT_BIG_MODEL; print(DEFAULT_BIG_MODEL)") + ONNX_REF=$(git log -1 --format='%H' -- openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx) + echo "model_name=${NAME}" >> $GITHUB_OUTPUT + echo "onnx_ref=$ONNX_REF" >> $GITHUB_OUTPUT + + build_model: + needs: resolve_name + uses: ./.github/workflows/sunnypilot-build-model.yaml + with: + upstream_branch: ${{ needs.resolve_name.outputs.onnx_ref }} + custom_name: ${{ needs.resolve_name.outputs.model_name }} + target_hardware: usbgpu + secrets: inherit + + upload_defaults: + needs: [ resolve_name, build_model ] + runs-on: ubuntu-24.04 + permissions: + id-token: write + contents: write + steps: + - uses: actions/checkout@v4 + with: + submodules: recursive + - run: git lfs pull -I "openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx" + + - name: Install huggingface_hub + run: pip install --upgrade "huggingface_hub>=0.22.0" + + - name: Download artifact name + uses: actions/download-artifact@v4 + with: + name: artifact-name-${{ needs.resolve_name.outputs.model_name }} + path: artifact_name + + - name: Read artifact name + id: artifact + run: | + ARTIFACT_NAME=$(cat artifact_name/artifact_name.txt) + echo "artifact_name=$ARTIFACT_NAME" >> $GITHUB_OUTPUT + + - name: Download model artifact + uses: actions/download-artifact@v4 + with: + name: ${{ steps.artifact.outputs.artifact_name }} + path: output + + - name: Upload to HF and update default_models.json + env: + HF_OIDC_RESOURCE: datasets/${{ env.HF_REPO }} + ARTIFACT_NAME: ${{ steps.artifact.outputs.artifact_name }} + run: | + rm -f output/artifact_name.txt + export PYTHONPATH=$(pwd) + python3 release/ci/upload_default_model.py \ + --hf-repo "${{ env.HF_REPO }}" \ + --hf-defaults-path "${{ env.HF_DEFAULTS_PATH }}" \ + --artifact-name "$ARTIFACT_NAME" \ + --model-dir output \ + --onnx-path "openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx" \ + --onnx-ref "${{ needs.resolve_name.outputs.onnx_ref }}" \ + --model-name "${{ needs.resolve_name.outputs.model_name }}" \ + --tinygrad-ref "$(python3 openpilot/sunnypilot/models/tinygrad_ref.py)" \ + --run-number "${{ github.run_number }}" diff --git a/.github/workflows/sunnypilot-build-prebuilt.yaml b/.github/workflows/sunnypilot-build-prebuilt.yaml index db0d0f55b4..e155a68d4a 100644 --- a/.github/workflows/sunnypilot-build-prebuilt.yaml +++ b/.github/workflows/sunnypilot-build-prebuilt.yaml @@ -36,6 +36,7 @@ jobs: publish_concurrency_group: ${{ steps.strategy.outputs.publish_concurrency_group }} is_stable_branch: ${{ steps.strategy.outputs.is_stable_branch }} build: ${{ steps.strategy.outputs.build }} + include_big_model: ${{ steps.strategy.outputs.include_big_model }} steps: - uses: actions/checkout@v4 - name: Extract deploy strategy @@ -78,6 +79,9 @@ jobs: stable_version=$(cat openpilot/sunnypilot/common/version.h | grep SUNNYPILOT_VERSION | sed -e 's/[^0-9|.]//g'); echo "version=$([ "$is_stable_branch" = "true" ] && echo "$stable_version" || echo "$BUILD")" >> $GITHUB_OUTPUT echo "extra_version_identifier=${environment}" >> $GITHUB_OUTPUT + + include_big_model="$(echo "$CONFIG" | jq -r '.include_big_model // false')"; + echo "include_big_model=$include_big_model" >> $GITHUB_OUTPUT fi echo "build=$BUILD" >> $GITHUB_OUTPUT cat $GITHUB_OUTPUT @@ -203,6 +207,101 @@ jobs: source ${UV_PROJECT_ENVIRONMENT}/bin/activate PYTHONPATH=$PYTHONPATH:${{ github.workspace }}/ ${{ github.workspace }}/scripts/manage-powersave.py --enable + prepare_chestnut: + needs: [ prepare_strategy ] + runs-on: ubuntu-24.04 + if: ${{ needs.prepare_strategy.outputs.include_big_model == 'true' }} + env: + HF_REPO: sunnypilot/sunnypilot_models_v1 + HF_DEFAULTS_PATH: models/defaults/big + steps: + - uses: actions/checkout@v4 + with: + ref: ${{ github.head_ref || github.ref_name }} + + - run: git lfs pull -I "openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx" + + - name: Check HF defaults and build if needed + id: resolve + run: | + ACTUAL_ONNX_HASH=$(sha256sum "openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx" | cut -d' ' -f1) + echo "Repo ONNX hash: $ACTUAL_ONNX_HASH" + + JSON_URL="https://huggingface.co/datasets/${HF_REPO}/resolve/main/${HF_DEFAULTS_PATH}/default_models.json" + + check_hash() { + DEFAULTS=$(curl -fsSL "$JSON_URL" 2>/dev/null) || return 1 + BUNDLE=$(echo "$DEFAULTS" | jq --arg hash "$ACTUAL_ONNX_HASH" '.bundles[] | select(.onnx_sha256 == $hash)' 2>/dev/null) + [ -n "$BUNDLE" ] && [ "$BUNDLE" != "null" ] + } + + if check_hash; then + echo "HF defaults match repo ONNX" + else + echo "No matching model on HF — triggering build" + gh workflow run build-default-big-model.yaml --ref "${{ github.head_ref || github.ref_name }}" + + echo "Waiting for build to start..." + sleep 120 + + RUN_ID=$(gh run list --workflow=build-default-big-model.yaml --branch="${{ github.head_ref || github.ref_name }}" --limit=1 --json databaseId --jq '.[0].databaseId') + if [ -z "$RUN_ID" ] || [ "$RUN_ID" = "null" ]; then + echo "::error::Failed to find build-default-big-model run" + exit 1 + fi + + echo "Waiting for run $RUN_ID..." + gh run watch "$RUN_ID" + + CONCLUSION=$(gh run view "$RUN_ID" --json conclusion --jq '.conclusion') + if [ "$CONCLUSION" != "success" ]; then + echo "::error::build-default-big-model failed: $CONCLUSION" + exit 1 + fi + + if ! check_hash; then + echo "::error::HF defaults still don't match after build" + exit 1 + fi + fi + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + + - name: Download big model chunks + run: | + ACTUAL_ONNX_HASH=$(sha256sum "openpilot/selfdrive/modeld/models/big_driving_supercombo.onnx" | cut -d' ' -f1) + JSON_URL="https://huggingface.co/datasets/${HF_REPO}/resolve/main/${HF_DEFAULTS_PATH}/default_models.json" + DEFAULTS=$(curl -fsSL "$JSON_URL") + BUNDLE=$(echo "$DEFAULTS" | jq --arg hash "$ACTUAL_ONNX_HASH" '.bundles[] | select(.onnx_sha256 == $hash)') + + mkdir -p big_model_chunks + ARTIFACT=$(echo "$BUNDLE" | jq -r '.models[0].artifact') + BASE_URL=$(echo "$ARTIFACT" | jq -r '.download_uri.url' | sed 's|/[^/]*$||') + NUM_CHUNKS=$(echo "$ARTIFACT" | jq -r '.chunks | length') + + CANONICAL="big_driving_tinygrad.pkl" + echo "$ARTIFACT" | jq -r '.chunks[].file_name' | while read CHUNK_NAME; do + CHUNK_IDX=$(echo "$CHUNK_NAME" | grep -oP 'chunk\K[0-9]+of[0-9]+') + CANONICAL_CHUNK="${CANONICAL}.chunk${CHUNK_IDX}" + ENCODED_URL=$(python3 -c "import urllib.parse; print(urllib.parse.quote('${BASE_URL}/${CHUNK_NAME}', safe=':/'))") + echo "Downloading $CHUNK_NAME -> $CANONICAL_CHUNK" + curl -fsSL -o "big_model_chunks/${CANONICAL_CHUNK}" "$ENCODED_URL" + done + + echo "$NUM_CHUNKS" > "big_model_chunks/${CANONICAL}.chunkmanifest" + + - name: Upload big model chunks + uses: actions/upload-artifact@v4 + with: + name: big-model-chunks + path: big_model_chunks/ + compression-level: 0 + + - name: Cancel run on failure + if: failure() + run: gh run cancel ${{ github.run_id }} + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} publish: concurrency: @@ -211,14 +310,20 @@ jobs: # Otherwise, if a job is waiting to be published due to environment wait time, it would be canceled by a new commit and restart the wait time. group: ${{ needs.prepare_strategy.outputs.publish_concurrency_group }} cancel-in-progress: ${{ needs.prepare_strategy.outputs.cancel_publish_in_progress == 'true' }} - if: ${{ (always() && !cancelled() && !failure()) && needs.build.result == 'success' && needs.prepare_strategy.result == 'success' && (!contains(github.event_name, 'pull_request') || (github.event.action == 'labeled' && github.event.label.name == 'prebuilt')) }} - needs: [ build, prepare_strategy ] + if: ${{ + always() && !cancelled() && + needs.build.result == 'success' && + needs.prepare_strategy.result == 'success' && + (!contains(github.event_name, 'pull_request') || (github.event.action == 'labeled' && github.event.label.name == 'prebuilt')) && + (needs.prepare_strategy.outputs.include_big_model != 'true' || needs.prepare_chestnut.result == 'success') + }} + needs: [ build, prepare_strategy, prepare_chestnut ] runs-on: ubuntu-24.04 environment: ${{ needs.prepare_strategy.outputs.environment }} steps: - uses: actions/checkout@v4 - - name: Download build artifacts + - name: Download prebuilt artifact uses: actions/download-artifact@v4 with: name: prebuilt @@ -228,6 +333,24 @@ jobs: mkdir -p ${{ env.OUTPUT_DIR }} tar xzf prebuilt.tar.gz -C ${{ env.OUTPUT_DIR }} + - name: Prepare chestnut output + if: ${{ needs.prepare_chestnut.result == 'success' }} + run: | + mkdir -p "${{ github.workspace }}/chestnut_output" + tar xzf prebuilt.tar.gz -C "${{ github.workspace }}/chestnut_output" + + - name: Download big model chunks + if: ${{ needs.prepare_chestnut.result == 'success' }} + uses: actions/download-artifact@v4 + with: + name: big-model-chunks + path: big_model_chunks + + - name: Inject big model into chestnut + if: ${{ needs.prepare_chestnut.result == 'success' }} + run: | + cp big_model_chunks/* "${{ github.workspace }}/chestnut_output/openpilot/selfdrive/modeld/models/" + - name: Configure Git run: | git config --global user.email "github-actions[bot]@users.noreply.github.com" @@ -248,6 +371,22 @@ jobs: "https://x-access-token:${{github.token}}@github.com/sunnypilot/sunnypilot.git" \ "${{ needs.prepare_strategy.outputs.extra_version_identifier }}" + - name: Publish chestnut branch + if: ${{ needs.prepare_chestnut.result == 'success' }} + env: + GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} + run: | + CHESTNUT_BRANCH="${{ needs.prepare_strategy.outputs.new_branch }}-chestnut" + CHESTNUT_DIR="${{ github.workspace }}/chestnut_output" + + ${{ env.CI_DIR }}/publish.sh \ + "${{ github.workspace }}" \ + "$CHESTNUT_DIR" \ + "$CHESTNUT_BRANCH" \ + "${{ needs.prepare_strategy.outputs.version }}" \ + "https://x-access-token:${{github.token}}@github.com/sunnypilot/sunnypilot.git" \ + "${{ needs.prepare_strategy.outputs.extra_version_identifier }}" + - name: Tag ${{ needs.prepare_strategy.outputs.environment }} if: ${{ needs.prepare_strategy.outputs.is_stable_branch == 'true' && (github.event_name != 'push' || !startsWith(github.ref, 'refs/tags/')) }} run: | @@ -260,6 +399,7 @@ jobs: - prepare_strategy - build - publish + - prepare_chestnut runs-on: ubuntu-24.04 if: ${{ (always() && !cancelled() && !failure()) && needs.publish.result == 'success' @@ -279,6 +419,7 @@ jobs: export commit_short_sha="${commit_short_sha:0:7}" export extra_version_identifier="${{ needs.prepare_strategy.outputs.extra_version_identifier || github.run_number }}" export PUBLIC_REPO_URL="${{ env.PUBLIC_REPO_URL }}" + export chestnut_branch="${{ needs.prepare_chestnut.result == 'success' && format('{0}-chestnut', needs.prepare_strategy.outputs.new_branch) || '' }}" MESSAGE=$(cat << 'EOF' | envsubst ${{ vars.DISCOURSE_GENERAL_UPDATE_NOTICE }} diff --git a/openpilot/common/version.py b/openpilot/common/version.py index f1514aa3cd..0456782c05 100755 --- a/openpilot/common/version.py +++ b/openpilot/common/version.py @@ -16,6 +16,15 @@ MASTER_SP_BRANCHES = ['master'] RELEASE_BRANCHES = ['release-tizi-staging', 'release-mici-staging', 'release-tizi', 'release-mici', 'nightly'] TESTED_BRANCHES = RELEASE_BRANCHES + ['devel-staging', 'nightly-dev'] + RELEASE_SP_BRANCHES + TESTED_SP_BRANCHES +CHESTNUT_BRANCHES = { + "staging": "staging-chestnut", + "dev": "dev-chestnut", + "release-mici": "release-chestnut", + "release-tizi": "release-chestnut", + "release-mici-staging": "release-chestnut-staging", + "release-tizi-staging": "release-chestnut-staging", +} + SP_BRANCH_MIGRATIONS = { ("tici", "staging-c3-new"): "staging-tici", ("tici", "dev-c3-new"): "staging-tici", diff --git a/openpilot/selfdrive/selfdrived/alerts_offroad.json b/openpilot/selfdrive/selfdrived/alerts_offroad.json index 226be06683..91a7ec8ab3 100644 --- a/openpilot/selfdrive/selfdrived/alerts_offroad.json +++ b/openpilot/selfdrive/selfdrived/alerts_offroad.json @@ -18,7 +18,7 @@ "_comment": "Set extra field to the failed reason." }, "Offroad_ChestnutBranch": { - "text": "Chestnut detected! Switch to the release-chestnut branch to use chestnut-class models.", + "text": "Chestnut detected! Switch to the %1 branch to use chestnut-class models.", "severity": 0 }, "Offroad_UnregisteredHardware": { diff --git a/openpilot/selfdrive/ui/onroad/model_renderer.py b/openpilot/selfdrive/ui/onroad/model_renderer.py index def1c644af..f5a39a2a5b 100644 --- a/openpilot/selfdrive/ui/onroad/model_renderer.py +++ b/openpilot/selfdrive/ui/onroad/model_renderer.py @@ -192,7 +192,7 @@ class ModelRenderer(Widget, ChevronMetrics, ModelRendererSP): max_idx = self._get_path_length_idx(path_x_array, max_distance) self._path.projected_points = self._map_line_to_polygon( - self._path.raw_points, 0.9, self._path_offset_z, max_idx, max_distance, allow_invert=False + self._path.raw_points, self._get_path_half_width(), self._path_offset_z, max_idx, max_distance, allow_invert=False ) self._update_experimental_gradient() @@ -292,7 +292,7 @@ class ModelRenderer(Widget, ChevronMetrics, ModelRendererSP): allow_throttle = sm['longitudinalPlan'].allowThrottle or not self._longitudinal_control self._blend_filter.update(int(allow_throttle)) - if ui_state.rainbow_path: + if ui_state.rainbow_path and self._lateral_active: self.rainbow_path.draw_rainbow_path(self._rect, self._path) return diff --git a/openpilot/selfdrive/ui/sunnypilot/onroad/model_renderer.py b/openpilot/selfdrive/ui/sunnypilot/onroad/model_renderer.py index 5d78997662..3cf639d0a1 100644 --- a/openpilot/selfdrive/ui/sunnypilot/onroad/model_renderer.py +++ b/openpilot/selfdrive/ui/sunnypilot/onroad/model_renderer.py @@ -4,11 +4,23 @@ 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 openpilot.common.filter_simple import FirstOrderFilter +from openpilot.selfdrive.ui.ui_state import ui_state, UIStatus from openpilot.selfdrive.ui.sunnypilot.onroad.chevron_metrics import ChevronMetrics from openpilot.selfdrive.ui.sunnypilot.onroad.rainbow_path import RainbowPath +from openpilot.system.ui.lib.application import gui_app class ModelRendererSP: def __init__(self): self.rainbow_path = RainbowPath() self.chevron_metrics = ChevronMetrics() + self._width_filter = FirstOrderFilter(0.9, 0.1, 1 / gui_app.target_fps) + + @property + def _lateral_active(self) -> bool: + return ui_state.status in (UIStatus.ENGAGED, UIStatus.LAT_ONLY) + + def _get_path_half_width(self) -> float: + target = 0.9 if self._lateral_active else 0.40 + return self._width_filter.update(target) diff --git a/openpilot/system/hardware/hardwared.py b/openpilot/system/hardware/hardwared.py index 3c22f1cbd7..8aed8be5b9 100755 --- a/openpilot/system/hardware/hardwared.py +++ b/openpilot/system/hardware/hardwared.py @@ -27,7 +27,7 @@ from openpilot.common.swaglog import cloudlog from openpilot.sunnypilot.system.statsd import statlog from openpilot.system.hardware.power_monitoring import PowerMonitoring from openpilot.system.hardware.fan_controller import FanController -from openpilot.common.version import terms_version, training_version, get_build_metadata, terms_version_sp +from openpilot.common.version import terms_version, training_version, get_build_metadata, terms_version_sp, CHESTNUT_BRANCHES ThermalStatus = log.DeviceState.ThermalStatus @@ -301,7 +301,11 @@ def hardware_thread(end_event, hw_queue) -> None: set_usb_state(msg.deviceState, last_hw_state.usb_state) chestnut.update(started_ts is None, last_hw_state.usb_state) - set_offroad_alert_if_changed("Offroad_ChestnutBranch", msg.deviceState.chestnutPresent and not big_model_available) + current_channel = get_build_metadata().channel + chestnut_target = CHESTNUT_BRANCHES.get(current_channel) + chestnut_needs_switch = msg.deviceState.chestnutPresent and not big_model_available and chestnut_target is not None + set_offroad_alert_if_changed("Offroad_ChestnutBranch", chestnut_needs_switch, + extra_text=chestnut_target if chestnut_needs_switch else None) # this subset is only used for offroad temp_sources = [ diff --git a/release/ci/model_generator.py b/release/ci/model_generator.py index 80b102a6f1..ff9be64783 100755 --- a/release/ci/model_generator.py +++ b/release/ci/model_generator.py @@ -90,6 +90,18 @@ def _rename_pkl_with_chunks(old_pkl: Path, new_pkl: Path) -> Path: return old_pkl.rename(new_pkl) +def _hash_onnx_files(model_dir: Path) -> str | None: + onnx_files = sorted(model_dir.glob("*.onnx")) + if not onnx_files: + return None + digest = hashlib.sha256() + for f in onnx_files: + with f.open('rb') as fh: + while block := fh.read(1024 * 1024): + digest.update(block) + return digest.hexdigest() + + def generate_chunked_model(driving_pkl: Path) -> dict: tinygrad_hash = _hash_pkl(driving_pkl) @@ -123,7 +135,8 @@ def generate_chunked_model(driving_pkl: Path) -> dict: } -def create_metadata_json(models: list, output_dir: Path, custom_name=None, short_name=None, is_20hz=False, upstream_branch="unknown") -> None: +def create_metadata_json(models: list, output_dir: Path, custom_name=None, short_name=None, is_20hz=False, upstream_branch="unknown", + onnx_sha256=None) -> None: bundle_json = { "short_name": short_name, "display_name": custom_name or upstream_branch, @@ -139,6 +152,9 @@ def create_metadata_json(models: list, output_dir: Path, custom_name=None, short "models": models, } + if onnx_sha256: + bundle_json["onnx_sha256"] = onnx_sha256 + # Write metadata to output_dir metadata_json = { "bundles": [bundle_json] @@ -178,4 +194,6 @@ if __name__ == "__main__": _driving_pkl = new_pkl _model_metadata = generate_chunked_model(_driving_pkl) - create_metadata_json([_model_metadata], _output_dir, args.custom_name, _short_name, args.is_20hz, args.upstream_branch) + _onnx_sha256 = _hash_onnx_files(Path(args.model_dir)) + create_metadata_json([_model_metadata], _output_dir, args.custom_name, _short_name, args.is_20hz, args.upstream_branch, + onnx_sha256=_onnx_sha256) diff --git a/release/ci/upload_default_model.py b/release/ci/upload_default_model.py new file mode 100644 index 0000000000..eac48e3be4 --- /dev/null +++ b/release/ci/upload_default_model.py @@ -0,0 +1,104 @@ +#!/usr/bin/env python3 +""" +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. +""" + +import argparse +import hashlib +import json +import tempfile + +from huggingface_hub import HfApi, hf_hub_download + + +def hash_file(path: str) -> str: + digest = hashlib.sha256() + with open(path, 'rb') as f: + while block := f.read(1024 * 1024): + digest.update(block) + return digest.hexdigest() + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--hf-repo", required=True) + parser.add_argument("--hf-defaults-path", required=True) + parser.add_argument("--artifact-name", required=True) + parser.add_argument("--model-dir", required=True) + parser.add_argument("--onnx-path", required=True) + parser.add_argument("--onnx-ref", required=True) + parser.add_argument("--model-name", required=True) + parser.add_argument("--tinygrad-ref", required=True) + parser.add_argument("--run-number", required=True) + args = parser.parse_args() + + api = HfApi() + onnx_sha256 = hash_file(args.onnx_path) + short_ref = args.onnx_ref[:8] + 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})") + print(f"Folder: {folder_name}") + + metadata_path = f"{args.model_dir}/metadata.json" + with open(metadata_path) as f: + metadata = json.load(f) + + bundle = metadata['bundles'][0] + bundle['display_name'] = args.model_name + bundle['onnx_sha256'] = onnx_sha256 + bundle['onnx_ref'] = args.onnx_ref + + artifact = bundle['models'][0]['artifact'] + hf_base = f"https://huggingface.co/datasets/{args.hf_repo}/resolve/main/{args.hf_defaults_path}/{folder_name}" + artifact['download_uri']['url'] = f"{hf_base}/{artifact['file_name']}" + for chunk in artifact.get('chunks', []): + chunk['url'] = f"{hf_base}/{chunk['file_name']}" + + print(f"Uploading model to {args.hf_defaults_path}/{folder_name}/") + api.upload_folder( + folder_path=args.model_dir, + path_in_repo=f"{args.hf_defaults_path}/{folder_name}", + repo_id=args.hf_repo, + repo_type="dataset", + ) + + json_filename = f"{args.hf_defaults_path}/default_models.json" + try: + local_path = hf_hub_download(repo_id=args.hf_repo, repo_type='dataset', filename=json_filename) + with open(local_path) as f: + defaults_json = json.load(f) + except Exception: + defaults_json = {"tinygrad_ref": args.tinygrad_ref, "bundles": []} + + defaults_json['tinygrad_ref'] = args.tinygrad_ref + + existing_idx = next((i for i, b in enumerate(defaults_json['bundles']) + if b.get('onnx_sha256') == onnx_sha256), None) + if existing_idx is not None: + defaults_json['bundles'][existing_idx] = bundle + else: + defaults_json['bundles'].append(bundle) + + print(json.dumps(defaults_json, indent=2)) + + with tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False) as f: + json.dump(defaults_json, f, indent=2) + tmp_path = f.name + + api.upload_file( + path_or_fileobj=tmp_path, + path_in_repo=json_filename, + repo_id=args.hf_repo, + repo_type="dataset", + ) + + print(f"Updated {json_filename}") + + +if __name__ == "__main__": + main()