mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-09-30 03:13:48 +08:00
Model Stuff
This commit is contained in:
+134
-22
@@ -29,6 +29,9 @@ from pathlib import Path
|
||||
|
||||
|
||||
OPENPILOT_REPO = "commaai/openpilot"
|
||||
# openpilot's .lfsconfig points LFS storage at Hugging Face, not github.com's own LFS.
|
||||
# The github.com LFS batch endpoint now 404s on every object; use the real backend.
|
||||
OPENPILOT_LFS_BATCH_URL = "https://huggingface.co/commaai/openpilot-lfs.git/info/lfs/objects/batch"
|
||||
RESOURCES_REPO = os.environ.get("STARPILOT_RESOURCES_REPO", "firestar5683/StarPilot-Resources")
|
||||
HF_BUCKET = os.environ.get("STARPILOT_HF_BUCKET", "StarPilot-Driving/StarPilot-Resources")
|
||||
RESOURCE_BRANCH = "Models"
|
||||
@@ -80,7 +83,8 @@ def default_workspace() -> Path:
|
||||
return Path.home() / "Desktop" / "StarPilot-Model-Releases"
|
||||
|
||||
|
||||
def run(command: list[str], *, cwd: Path | None = None, capture: bool = False) -> subprocess.CompletedProcess:
|
||||
def run(command: list[str], *, cwd: Path | None = None, capture: bool = False,
|
||||
env: dict[str, str] | None = None) -> subprocess.CompletedProcess:
|
||||
print("$ " + " ".join(shlex.quote(part) for part in command))
|
||||
return subprocess.run(
|
||||
command,
|
||||
@@ -88,6 +92,7 @@ def run(command: list[str], *, cwd: Path | None = None, capture: bool = False) -
|
||||
check=True,
|
||||
text=capture,
|
||||
capture_output=capture,
|
||||
env=env,
|
||||
)
|
||||
|
||||
|
||||
@@ -160,7 +165,6 @@ def stream_response(response, destination: Path, prefix: bytes = b"") -> tuple[i
|
||||
|
||||
|
||||
def download_lfs_object(oid: str, expected_size: int, ref: str, destination: Path) -> tuple[int, str]:
|
||||
batch_url = f"https://github.com/{OPENPILOT_REPO}.git/info/lfs/objects/batch"
|
||||
payload = json.dumps({
|
||||
"operation": "download",
|
||||
"transfers": ["basic"],
|
||||
@@ -168,7 +172,7 @@ def download_lfs_object(oid: str, expected_size: int, ref: str, destination: Pat
|
||||
"ref": {"name": ref},
|
||||
}).encode("utf-8")
|
||||
with http_request(
|
||||
batch_url,
|
||||
OPENPILOT_LFS_BATCH_URL,
|
||||
method="POST",
|
||||
payload=payload,
|
||||
headers={"Accept": "application/vnd.git-lfs+json", "Content-Type": "application/vnd.git-lfs+json"},
|
||||
@@ -395,6 +399,38 @@ def parse_pasted_release(text: str, model_id_override: str | None, behavior_vers
|
||||
)
|
||||
|
||||
|
||||
def release_from_source_file(source_file: Path, model_id_override: str | None, display_name_override: str | None,
|
||||
behavior_version: str) -> ReleaseInfo:
|
||||
"""Build a ReleaseInfo for a local ONNX that didn't come from an openpilot git commit.
|
||||
|
||||
Used when the model bot / commit only ships a precompiled artifact (e.g. an eGPU pkl)
|
||||
and the real ONNX source has to be fetched by hand from elsewhere, such as
|
||||
huggingface.co/commaai/openpilot_driving_models.
|
||||
"""
|
||||
if not source_file.is_file():
|
||||
raise ReleaseError(f"--source-file not found: {source_file}")
|
||||
if source_file.suffix.lower() != ".onnx":
|
||||
raise ReleaseError(f"--source-file must be a .onnx file: {source_file}")
|
||||
display_name = display_name_override or source_name_from_path(source_file.name)
|
||||
model_id = model_id_override or slug_model_id(display_name)
|
||||
if not MODEL_ID_RE.fullmatch(model_id):
|
||||
raise ReleaseError(f"Invalid model ID {model_id!r}; use lowercase letters, digits, '-' or '_' (pass --model-id)")
|
||||
iteration_match = re.search(r"\b(v\d+)\b", display_name, flags=re.IGNORECASE)
|
||||
return ReleaseInfo(
|
||||
model_id=model_id,
|
||||
display_name=display_name,
|
||||
release_date=dt.date.today().isoformat(),
|
||||
branch="",
|
||||
source_ref="",
|
||||
source_path=str(source_file),
|
||||
input_format="supercombo" if "supercombo" in source_file.name.lower() else "split",
|
||||
behavior_version=behavior_version,
|
||||
uses_external_gpu=source_file.name.lower().startswith("big_"),
|
||||
commits=[],
|
||||
model_iteration=iteration_match.group(1).lower() if iteration_match else "",
|
||||
)
|
||||
|
||||
|
||||
def resolve_branch_commit(branch: str) -> str:
|
||||
url = f"https://api.github.com/repos/{OPENPILOT_REPO}/commits/{urllib.parse.quote(branch, safe='')}"
|
||||
payload = get_json(url)
|
||||
@@ -508,11 +544,16 @@ def remote_compile(info: ReleaseInfo, source: Path, ip: str, workspace: Path, ke
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
text=True,
|
||||
encoding="utf-8",
|
||||
errors="replace",
|
||||
bufsize=1,
|
||||
)
|
||||
assert process.stdout is not None
|
||||
for line in process.stdout:
|
||||
print(f"[device] {line}", end="")
|
||||
try:
|
||||
print(f"[device] {line}", end="")
|
||||
except UnicodeEncodeError:
|
||||
print(f"[device] {line}".encode(sys.stdout.encoding or 'utf-8', errors='replace').decode(sys.stdout.encoding or 'utf-8'), end="")
|
||||
log.write(line)
|
||||
return_code = process.wait()
|
||||
if return_code != 0:
|
||||
@@ -684,18 +725,62 @@ def validate_manifest_update(before: object, after: object, replacing_model_id:
|
||||
)
|
||||
|
||||
|
||||
def find_hf() -> str:
|
||||
candidates = [shutil.which("hf"), str(Path.home() / ".local/bin/hf")]
|
||||
for candidate in candidates:
|
||||
def windows_user_site_candidates(python: str) -> list[str]:
|
||||
"""Compute the per-user site-packages path Windows Python installs use.
|
||||
|
||||
On some Windows installs, a Python interpreter that has huggingface_hub
|
||||
importable interactively still fails to find it as a subprocess child: both
|
||||
`import huggingface_hub` and even `pip show` fail the same way in that
|
||||
child, because site.getusersitepackages() resolves to the wrong per-app
|
||||
path (a Windows Store Python aliasing quirk) instead of that interpreter's
|
||||
own user site-packages. That makes any subprocess-based probe of the
|
||||
interpreter circular, so compute the conventional path directly instead:
|
||||
%APPDATA%\\Python\\PythonXY\\site-packages. Querying sys.version_info alone
|
||||
(no package import) does not trigger the broken resolution.
|
||||
"""
|
||||
appdata = os.environ.get("APPDATA")
|
||||
if not appdata:
|
||||
return []
|
||||
version = subprocess.run(
|
||||
[python, "-c", "import sys; print(f'{sys.version_info[0]}{sys.version_info[1]}')"],
|
||||
capture_output=True, text=True,
|
||||
)
|
||||
if version.returncode != 0 or not version.stdout.strip().isdigit():
|
||||
return []
|
||||
return [os.path.join(appdata, "Python", f"Python{version.stdout.strip()}", "site-packages")]
|
||||
|
||||
|
||||
def find_hf() -> tuple[list[str], dict[str, str] | None]:
|
||||
"""Return (command prefix, extra env) to invoke the Hugging Face CLI.
|
||||
|
||||
The installed `hf` launcher executable can fail to resolve its own Python
|
||||
environment when invoked directly via subprocess (no shell/PATH context),
|
||||
even though it works fine from an interactive shell. Prefer running the CLI
|
||||
as a module of whichever Python actually has huggingface_hub importable.
|
||||
"""
|
||||
for python in (sys.executable, shutil.which("python3"), shutil.which("python")):
|
||||
if not python:
|
||||
continue
|
||||
probe = subprocess.run([python, "-c", "import huggingface_hub"], capture_output=True)
|
||||
if probe.returncode == 0:
|
||||
return [python, "-m", "huggingface_hub.cli.hf"], None
|
||||
for site_packages in windows_user_site_candidates(python):
|
||||
if not os.path.isdir(os.path.join(site_packages, "huggingface_hub")):
|
||||
continue
|
||||
env = {**os.environ, "PYTHONPATH": os.pathsep.join([site_packages, os.environ.get("PYTHONPATH", "")])}
|
||||
probe = subprocess.run([python, "-c", "import huggingface_hub"], capture_output=True, env=env)
|
||||
if probe.returncode == 0:
|
||||
return [python, "-m", "huggingface_hub.cli.hf"], env
|
||||
for candidate in (shutil.which("hf"), str(Path.home() / ".local/bin/hf")):
|
||||
if candidate and Path(candidate).is_file():
|
||||
return candidate
|
||||
return [candidate], None
|
||||
raise ReleaseError("Hugging Face CLI not found; install/authenticate `hf` first")
|
||||
|
||||
|
||||
def hf_copy(source: Path, bucket: str, remote_path: str) -> None:
|
||||
hf = find_hf()
|
||||
command, env = find_hf()
|
||||
destination = f"hf://buckets/{bucket}/{remote_path}"
|
||||
run([hf, "buckets", "cp", str(source), destination, "--format", "quiet"])
|
||||
run([*command, "buckets", "cp", str(source), destination, "--format", "quiet"], env=env)
|
||||
|
||||
|
||||
def refresh_huggingface_manifest(manifest: Path, bucket: str) -> dict:
|
||||
@@ -704,7 +789,8 @@ def refresh_huggingface_manifest(manifest: Path, bucket: str) -> dict:
|
||||
manifest.parent.mkdir(parents=True, exist_ok=True)
|
||||
with tempfile.TemporaryDirectory(prefix=".model-release-", dir=manifest.parent) as temporary_dir:
|
||||
candidate = Path(temporary_dir) / manifest.name
|
||||
run([find_hf(), "buckets", "cp", source, str(candidate), "--format", "quiet"])
|
||||
hf_command, hf_env = find_hf()
|
||||
run([*hf_command, "buckets", "cp", source, str(candidate), "--format", "quiet"], env=hf_env)
|
||||
try:
|
||||
payload = json.loads(candidate.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as error:
|
||||
@@ -812,6 +898,13 @@ def parse_args() -> argparse.Namespace:
|
||||
parser.add_argument("commit", nargs="?", help="A single openpilot commit SHA; metadata is resolved from GitHub.")
|
||||
parser.add_argument("--text", help="Release text, otherwise paste it into stdin.")
|
||||
parser.add_argument("--text-file", type=Path, help="Read the pasted release text from a file.")
|
||||
parser.add_argument("--source-file", type=Path,
|
||||
help="Use this local .onnx instead of downloading one from an openpilot commit/branch. "
|
||||
"For releases whose commit only ships a precompiled artifact (e.g. an eGPU pkl) "
|
||||
"and the real ONNX must be fetched by hand from elsewhere. Skips the runtime-change "
|
||||
"scan since there is no commit history to scan; review the source yourself first. "
|
||||
"Requires --model-id or a v-numbered name in --display-name.")
|
||||
parser.add_argument("--display-name", help="Display name to record in the manifest when using --source-file.")
|
||||
parser.add_argument("--model-id", help="Override the ID parsed from the model name.")
|
||||
parser.add_argument("--behavior-version", default=DEFAULT_BEHAVIOR_VERSION, help="Runtime behavior version (default: v16).")
|
||||
parser.add_argument("--ip", help="Comma IP; prompted interactively when omitted.")
|
||||
@@ -837,21 +930,32 @@ def main() -> int:
|
||||
try:
|
||||
if not re.fullmatch(r"v\d+", args.behavior_version.strip(), flags=re.IGNORECASE):
|
||||
raise ReleaseError("--behavior-version must look like v16")
|
||||
text = read_release_text(args)
|
||||
if not text.strip():
|
||||
raise ReleaseError("No release text was supplied")
|
||||
info = parse_pasted_release(text, args.model_id, args.behavior_version.strip().lower())
|
||||
|
||||
using_source_file = args.source_file is not None
|
||||
text = ""
|
||||
if using_source_file:
|
||||
info = release_from_source_file(args.source_file.expanduser().resolve(), args.model_id,
|
||||
args.display_name, args.behavior_version.strip().lower())
|
||||
else:
|
||||
text = read_release_text(args)
|
||||
if not text.strip():
|
||||
raise ReleaseError("No release text was supplied")
|
||||
info = parse_pasted_release(text, args.model_id, args.behavior_version.strip().lower())
|
||||
if args.gpu is not None:
|
||||
info.uses_external_gpu = args.gpu
|
||||
print_summary(info)
|
||||
|
||||
findings = scan_runtime_changes(info)
|
||||
if findings:
|
||||
print_runtime_warning(findings)
|
||||
if not args.allow_runtime_changes:
|
||||
return 2
|
||||
if using_source_file:
|
||||
print("\nRuntime scan skipped: --source-file has no commit history to scan.")
|
||||
print("Confirm yourself that no tinygrad/modeld/Chestnut runtime changes are needed before proceeding.")
|
||||
else:
|
||||
print("Runtime scan: no tinygrad/modeld runtime files changed in the supplied commits.")
|
||||
findings = scan_runtime_changes(info)
|
||||
if findings:
|
||||
print_runtime_warning(findings)
|
||||
if not args.allow_runtime_changes:
|
||||
return 2
|
||||
else:
|
||||
print("Runtime scan: no tinygrad/modeld runtime files changed in the supplied commits.")
|
||||
|
||||
if args.dry_run:
|
||||
print("Dry run complete; no device or repository changes made.")
|
||||
@@ -862,7 +966,15 @@ def main() -> int:
|
||||
for relative in ("onnx", "compiled", "logs", "results"):
|
||||
(workspace / relative).mkdir(parents=True, exist_ok=True)
|
||||
source = workspace / "onnx" / f"{info.model_id}_driving_supercombo.onnx"
|
||||
source_result = download_source(info.source_ref, info.source_path, source, args.force)
|
||||
if using_source_file:
|
||||
local_source = Path(info.source_path)
|
||||
if source.exists() and not args.force:
|
||||
raise ReleaseError(f"Source already exists: {source}; use --force to replace it")
|
||||
shutil.copy2(local_source, source)
|
||||
source_result = {"path": str(source), "size": source.stat().st_size, "sha256": sha256_file(source),
|
||||
"url": str(local_source), "ref": "", "git_path": ""}
|
||||
else:
|
||||
source_result = download_source(info.source_ref, info.source_path, source, args.force)
|
||||
(workspace / "release.txt").write_text(text, encoding="utf-8")
|
||||
(workspace / "source.json").write_text(json.dumps({**source_result, "model": info.__dict__}, indent=2) + "\n", encoding="utf-8")
|
||||
|
||||
|
||||
@@ -467,7 +467,54 @@ def make_run_supercombo(model_runner, metadata, frame_skip, image_history_pipeli
|
||||
return run_policy
|
||||
|
||||
|
||||
def compile_jit(jit, make_random_inputs, input_keys, make_queues):
|
||||
def stateful_image_shapes(metadata):
|
||||
shape = tuple(metadata['input_shapes']['new_img'])
|
||||
if len(shape) != 4 or shape[:2] != (2, 6):
|
||||
raise ValueError(f"Unsupported stateful image shape: {shape}")
|
||||
return dict.fromkeys(('img', 'big_img'), (1, *shape[1:]))
|
||||
|
||||
|
||||
def stateful_host_shapes(metadata):
|
||||
return {name: shape for name, shape in metadata['input_shapes'].items()
|
||||
if name != 'new_img' and name not in metadata['state_pairs']}
|
||||
|
||||
|
||||
def make_stateful_input_queues(metadata, device):
|
||||
queues, npy = make_warp_input_queues(stateful_image_shapes(metadata), 1, device)
|
||||
shapes = stateful_host_shapes(metadata)
|
||||
sizes = [math.prod(shape) for shape in shapes.values()]
|
||||
packed = np.zeros(sum(sizes), dtype=np.float32)
|
||||
npy.update({name: value.reshape(shape) for (name, shape), value in
|
||||
zip(shapes.items(), np.split(packed, np.cumsum(sizes[:-1])), strict=True)})
|
||||
queues['packed_npy_inputs'] = Tensor(packed, device='NPY').realize()
|
||||
for name in metadata['state_pairs']:
|
||||
queues[name] = Tensor(np.zeros(metadata['input_shapes'][name], dtype=metadata['input_dtypes'][name]),
|
||||
device=device).contiguous().realize()
|
||||
return queues, npy
|
||||
|
||||
|
||||
def make_run_stateful_supercombo(model_runner, metadata):
|
||||
shapes = stateful_host_shapes(metadata)
|
||||
sizes = [math.prod(shape) for shape in shapes.values()]
|
||||
|
||||
def run_policy(warped, packed_npy_inputs, **state):
|
||||
packed = packed_npy_inputs.to(Device.DEFAULT).realize()
|
||||
inputs = {name: value.reshape(shape).cast(model_runner.graph_inputs[name].dtype)
|
||||
for (name, shape), value in zip(shapes.items(), packed.split(sizes), strict=True)}
|
||||
inputs['new_img'] = warped.to(Device.DEFAULT).cast(model_runner.graph_inputs['new_img'].dtype)
|
||||
outputs = {name: value.contiguous() for name, value in model_runner(inputs | state).items()}
|
||||
for name, next_name in metadata['state_pairs'].items():
|
||||
if outputs[next_name].dtype != state[name].dtype:
|
||||
raise ValueError(f'State dtype mismatch: {name} -> {next_name}')
|
||||
# All reads of the previous state must finish before updating any history buffer.
|
||||
Tensor.realize(*outputs.values())
|
||||
Tensor.realize(*(state[name].assign(outputs[next_name]) for name, next_name in metadata['state_pairs'].items()))
|
||||
return outputs['outputs'].cast('float32'),
|
||||
|
||||
return run_policy
|
||||
|
||||
|
||||
def compile_jit(jit, make_random_inputs, input_keys, make_queues, validation_runs=1):
|
||||
seed = 42
|
||||
|
||||
def random_inputs_run(fn, current_seed, test_values=None, test_buffers=None, expect_match=True):
|
||||
@@ -475,7 +522,8 @@ def compile_jit(jit, make_random_inputs, input_keys, make_queues):
|
||||
np.random.seed(current_seed)
|
||||
Tensor.manual_seed(current_seed)
|
||||
testing = test_values is not None or test_buffers is not None
|
||||
run_count = 1 if testing else 3
|
||||
run_count = validation_runs if testing else max(3, validation_runs)
|
||||
values, buffers = [], []
|
||||
|
||||
for index in range(run_count):
|
||||
for value in npy.values():
|
||||
@@ -489,9 +537,9 @@ def compile_jit(jit, make_random_inputs, input_keys, make_queues):
|
||||
end = time.perf_counter()
|
||||
print(f" [{index + 1}/{run_count}] enqueue {(mid - start) * 1e3:6.2f} ms -- total {(end - start) * 1e3:6.2f} ms")
|
||||
|
||||
if index == 0:
|
||||
values = [np.copy(value.numpy()) for value in outputs]
|
||||
buffers = [np.copy(value.numpy()) for value in input_queues.values()]
|
||||
if index < validation_runs:
|
||||
values.extend(np.copy(value.numpy()) for value in outputs)
|
||||
buffers.extend(np.copy(value.numpy()) for value in input_queues.values())
|
||||
if not all(np.isfinite(value).all() for value in values):
|
||||
raise ValueError("Compiled JIT produced non-finite outputs")
|
||||
|
||||
@@ -602,13 +650,31 @@ def main():
|
||||
output["metadata"]["model"] = make_metadata_dict(model_path)
|
||||
validate_metadata(output["metadata"]["model"])
|
||||
policy_shapes = output["metadata"]["model"]["input_shapes"]
|
||||
frame_skip = args.frame_skip or derive_frame_skip(policy_shapes)
|
||||
make_policy_queues = partial(make_supercombo_input_queues, policy_shapes, frame_skip)
|
||||
run_policy = make_run_supercombo(
|
||||
model_runner, output["metadata"], frame_skip, args.image_history_pipeline,
|
||||
)
|
||||
image_shapes = policy_shapes
|
||||
policy_input_keys = FAST_POLICY_INPUTS if args.image_history_pipeline == IMAGE_HISTORY_IN_POLICY else SUPERCOMBO_POLICY_INPUTS
|
||||
if 'new_img' in policy_shapes:
|
||||
if args.image_history_pipeline != IMAGE_HISTORY_IN_POLICY:
|
||||
parser.error('ONNX-managed history requires --image-history-pipeline policy')
|
||||
metadata = output['metadata']['model']
|
||||
metadata['state_pairs'] = {name: f'next_{name}' for name in policy_shapes
|
||||
if f'next_{name}' in metadata['output_shapes']}
|
||||
if not metadata['state_pairs']:
|
||||
raise ValueError('Stateful supercombo is missing next-state outputs')
|
||||
metadata['input_dtypes'] = {name: np.dtype(spec.dtype.fmt).name for name, spec in model_runner.graph_inputs.items()}
|
||||
for name, next_name in metadata['state_pairs'].items():
|
||||
if policy_shapes[name] != metadata['output_shapes'][next_name]:
|
||||
raise ValueError(f'State shape mismatch: {name} -> {next_name}')
|
||||
frame_skip = 1
|
||||
make_policy_queues = partial(make_stateful_input_queues, metadata)
|
||||
run_policy = make_run_stateful_supercombo(model_runner, metadata)
|
||||
image_shapes = stateful_image_shapes(metadata)
|
||||
policy_input_keys = ('packed_npy_inputs', *metadata['state_pairs'])
|
||||
else:
|
||||
frame_skip = args.frame_skip or derive_frame_skip(policy_shapes)
|
||||
make_policy_queues = partial(make_supercombo_input_queues, policy_shapes, frame_skip)
|
||||
run_policy = make_run_supercombo(
|
||||
model_runner, output["metadata"], frame_skip, args.image_history_pipeline,
|
||||
)
|
||||
image_shapes = policy_shapes
|
||||
policy_input_keys = FAST_POLICY_INPUTS if args.image_history_pipeline == IMAGE_HISTORY_IN_POLICY else SUPERCOMBO_POLICY_INPUTS
|
||||
else:
|
||||
if not args.vision_onnx:
|
||||
parser.error("--vision-onnx is required for split models")
|
||||
@@ -675,6 +741,7 @@ def main():
|
||||
)
|
||||
output["run_policy"] = compile_jit(
|
||||
run_policy_jit, make_random_model_inputs, policy_input_keys, make_policy_queues,
|
||||
validation_runs=5 if output['metadata'].get('model', {}).get('state_pairs') else 1,
|
||||
)
|
||||
|
||||
model_w, model_h = args.model_size
|
||||
|
||||
@@ -49,6 +49,9 @@ from openpilot.selfdrive.modeld.compile_modeld import (
|
||||
derive_frame_skip,
|
||||
make_split_input_queues,
|
||||
make_supercombo_input_queues,
|
||||
make_stateful_input_queues,
|
||||
stateful_host_shapes,
|
||||
stateful_image_shapes,
|
||||
)
|
||||
from openpilot.selfdrive.modeld.helpers import get_tg_input_devices, load_oob, tinygrad_dev_config, usbgpu_present
|
||||
from openpilot.selfdrive.modeld.usbgpu_link import wait_usbgpu_link
|
||||
@@ -614,8 +617,15 @@ class ModelState:
|
||||
self.run_policy = artifact["run_policy"]
|
||||
self.warp_enqueue = artifact[(cam_w, cam_h)]
|
||||
self.can_prepare_only = self.image_history_pipeline == IMAGE_HISTORY_IN_WARP
|
||||
self.onnx_history = self.model_type == 'supercombo' and bool(self.metadata['model'].get('state_pairs'))
|
||||
|
||||
if self.model_type == "supercombo":
|
||||
if self.onnx_history:
|
||||
metadata = self.metadata['model']
|
||||
input_shapes = stateful_image_shapes(metadata)
|
||||
self.output_slices = metadata['output_slices']
|
||||
self.input_queues, self.npy = make_stateful_input_queues(metadata, self.QUEUE_DEV)
|
||||
self.policy_input_shapes = stateful_host_shapes(metadata)
|
||||
elif self.model_type == "supercombo":
|
||||
input_shapes = self.metadata["model"]["input_shapes"]
|
||||
self.output_slices = self.metadata["model"]["output_slices"]
|
||||
self.input_queues, self.npy = make_supercombo_input_queues(input_shapes, self.frame_skip, self.QUEUE_DEV)
|
||||
@@ -725,7 +735,9 @@ class ModelState:
|
||||
return parsed
|
||||
|
||||
def _reset_state(self) -> None:
|
||||
if self.model_type == "supercombo":
|
||||
if self.onnx_history:
|
||||
self.input_queues, self.npy = make_stateful_input_queues(self.metadata['model'], self.QUEUE_DEV)
|
||||
elif self.model_type == "supercombo":
|
||||
self.input_queues, self.npy = make_supercombo_input_queues(
|
||||
self.policy_input_shapes, self.frame_skip, self.QUEUE_DEV,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user