YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

Tool-selection bi-encoders (best seed @ best layer)

One bi-encoder per (model_id, dataset_id), stored under {model_id}__{dataset_id}/: bi_encoder.safetensors (projection heads) + config.json (everything else, including the scalar log_temp). manifest.json lists all configs.

A bi-encoder scores a base model's hidden state at best_layer (end-of-user position) against candidate tool keys. Tool keys are the base model's hidden states at tool_layer (24), token-mean-pooled, over this canonical text:

Function name: {name}
Description: {description}
Parameters: {comma-separated property keys, or (none)}

Load + run inference

import os, json, torch, numpy as np, torch.nn as nn, torch.nn.functional as F
from safetensors.torch import load_file
from huggingface_hub import snapshot_download
from transformers import AutoModelForCausalLM, AutoTokenizer

class BiEncoder(nn.Module):
    def __init__(self, hidden_dim, tool_dim, proj_dim, temp_init):
        super().__init__()
        self.h_proj = nn.Linear(hidden_dim, proj_dim)
        self.t_proj = nn.Linear(tool_dim, proj_dim)
        self.log_temp = nn.Parameter(torch.tensor(np.log(1.0/temp_init), dtype=torch.float32))
    def encode_q(self, h): return F.normalize(self.h_proj(h), dim=-1)
    def encode_k(self, t): return F.normalize(self.t_proj(t), dim=-1)

def tool_to_canonical_text(t):
    p = t.get("parameters", {})
    keys = (list(p.get("properties", {}).keys()) if isinstance(p, dict) and "properties" in p
            else (list(p.keys()) if isinstance(p, dict) else []))
    return (f"Function name: {t.get('name','')}\n"
            f"Description: {t.get('description','')}\n"
            f"Parameters: {', '.join(keys) if keys else '(none)'}")

def load_bi_encoder(repo_id, model_id, dataset_id, token=None):
    snap = snapshot_download(repo_id, repo_type="model", token=token)
    d = os.path.join(snap, f"{model_id}__{dataset_id}")
    cfg = json.load(open(os.path.join(d, "config.json")))
    sd = load_file(os.path.join(d, "bi_encoder.safetensors"))
    m = BiEncoder(cfg["hidden_dim"], cfg["tool_dim"], cfg["proj_dim"], cfg["temp_init"])
    m.load_state_dict({**sd, "log_temp": torch.tensor(float(cfg["log_temp"]))})
    m.eval()
    return m, cfg

# Hint for one example: base-model L{best_layer} hidden state at end-of-user (query_h),
# scored vs candidate tool keys encoded at L{tool_layer}, mean-pooled.
@torch.no_grad()
def predict_tool(m, cfg, query_h, candidate_tools, base_mdl, base_tok):
    keys = []
    for t in candidate_tools:
        inp = base_tok(tool_to_canonical_text(t), return_tensors="pt",
                       max_length=cfg["tool_max_tokens"], truncation=True).to(base_mdl.device)
        h = base_mdl(**inp, output_hidden_states=True).hidden_states[cfg["tool_layer"]][0]
        keys.append(h.float().mean(0))
    q = m.encode_q(query_h.float()); k = m.encode_k(torch.stack(keys))
    scores = q @ k.T
    return int(scores.argmax().item())
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support