Saluki / model /saluki.py
wuxing0105's picture
Upload folder using huggingface_hub
4d43445 verified
Raw
History Blame Contribute Delete
1.96 kB
"""Small public wrapper around the bundled Basenji Saluki implementation."""
from __future__ import annotations
import json
import os
from pathlib import Path
from typing import Any, Mapping
def load_params(path: str | Path) -> dict[str, Any]:
"""Load parameters and normalize historical scalar multi-head settings."""
with Path(path).open(encoding="utf-8") as handle:
params = json.load(handle)
model_params = dict(params["model"])
num_targets = model_params.get("num_targets", 1)
heads = int(model_params.get("heads", 1))
if not isinstance(num_targets, list):
model_params["num_targets"] = [num_targets] * heads
normalized = dict(params)
normalized["model"] = model_params
return normalized
class SalukiModel:
"""Construct, restore and run one head of the official Saluki network."""
def __init__(self, params: Mapping[str, Any], head: int = 0):
os.environ.setdefault("TF_USE_LEGACY_KERAS", "1")
from .basenji.rnann import RnaNN
model_params = dict(params["model"] if "model" in params else params)
num_targets = model_params.get("num_targets", 1)
heads = int(model_params.get("heads", 1))
if not isinstance(num_targets, list):
model_params["num_targets"] = [num_targets] * heads
self.network = RnaNN(model_params)
if not 0 <= head < len(self.network.models):
raise ValueError(
f"head must be in [0, {len(self.network.models) - 1}], got {head}"
)
self.head = head
self.network.model = self.network.models[head]
@property
def keras_model(self):
return self.network.models[self.head]
def restore(self, weight_path: str | Path) -> "SalukiModel":
self.network.restore(str(weight_path), head_i=self.head)
return self
def predict(self, inputs, **kwargs):
return self.keras_model.predict(inputs, **kwargs)