Fix offline GPU firmware and Connect streaming

This commit is contained in:
firestar5683
2026-09-06 21:18:04 -05:00
parent 67bf3adfd4
commit bb284d0bc4
17 changed files with 422 additions and 10 deletions
+55 -8
View File
@@ -1,8 +1,10 @@
#!/usr/bin/env python3
import argparse
import atexit
import hashlib
import math
import os
import pathlib
import pickle
import shutil
import tempfile
@@ -13,10 +15,26 @@ from functools import partial
import numpy as np
def _patch_tinygrad_fetch_fw():
import hashlib
import pathlib
FIRMWARE_ROOTS = (
pathlib.Path(__file__).resolve().parent / "firmware",
pathlib.Path("/lib/firmware"),
)
FIRMWARE_CACHE_DIR = pathlib.Path("/data/tinygrad_fw_cache")
def _read_firmware(path, sha256, compressed=False, zstandard_module=None):
if not path.is_file():
return None
blob = path.read_bytes()
if compressed:
if zstandard_module is None:
import zstandard as zstandard_module
blob = zstandard_module.ZstdDecompressor().stream_reader(blob).read()
return blob if hashlib.sha256(blob).hexdigest() == sha256 else None
def _patch_tinygrad_fetch_fw():
try:
import zstandard
except ImportError:
@@ -28,12 +46,41 @@ def _patch_tinygrad_fetch_fw():
return
def fetch_fw(path, name, sha256):
firmware_path = pathlib.Path(f"/lib/firmware/{path}/{name}.zst")
if firmware_path.is_file():
blob = zstandard.ZstdDecompressor().stream_reader(firmware_path.read_bytes()).read()
if hashlib.sha256(blob).hexdigest() == sha256:
for root in FIRMWARE_ROOTS:
if (blob := _read_firmware(root / path / f"{name}.zst", sha256, compressed=True, zstandard_module=zstandard)) is not None:
return blob
return original_fetch_fw(path, name, sha256)
cached_path = FIRMWARE_CACHE_DIR / path / f"{name}.{sha256}"
if (blob := _read_firmware(cached_path, sha256)) is not None:
return blob
last_error = None
for attempt in range(3):
if attempt:
time.sleep(5)
try:
blob = original_fetch_fw(path, name, sha256)
break
except Exception as error:
last_error = error
else:
assert last_error is not None
raise last_error
cache_path = None
try:
cached_path.parent.mkdir(parents=True, exist_ok=True)
with tempfile.NamedTemporaryFile(dir=cached_path.parent, delete=False) as cache_file:
cache_file.write(blob)
cache_path = pathlib.Path(cache_file.name)
cache_path.replace(cached_path)
except OSError:
try:
if cache_path is not None:
cache_path.unlink(missing_ok=True)
except OSError:
pass
return blob
helpers.fetch_fw = fetch_fw
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,27 @@
#!/usr/bin/env python3
import hashlib
from pathlib import Path
import pytest
import zstandard
from tinygrad.runtime.autogen.am import fw
CHESTNUT_FIRMWARE = (
"gc_12_0_0_imu.bin",
"gc_12_0_0_me.bin",
"gc_12_0_0_mec.bin",
"gc_12_0_0_pfp.bin",
"gc_12_0_0_rlc.bin",
"psp_14_0_2_sos.bin",
"sdma_7_0_0.bin",
"smu_14_0_2.bin",
)
@pytest.mark.parametrize("name", CHESTNUT_FIRMWARE)
def test_bundled_chestnut_firmware(name):
path = Path(__file__).parents[1] / "firmware" / "amdgpu" / f"{name}.zst"
blob = zstandard.ZstdDecompressor().stream_reader(path.read_bytes()).read()
assert hashlib.sha256(blob).hexdigest() == fw.hashes[name]