[sgl-router] Share model-file discovery for chat formatters (#39485)
This commit is contained in:
@@ -58,7 +58,6 @@ fn download_tokenizer_json(repo_id: &str) -> Result<std::path::PathBuf> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Download `file` from a HuggingFace repo id and return the cached local path.
|
/// Download `file` from a HuggingFace repo id and return the cached local path.
|
||||||
/// Shared by `tokenizer.json` (required) and `tokenizer_config.json` (optional).
|
|
||||||
fn download_repo_file(repo_id: &str, file: &str) -> Result<std::path::PathBuf> {
|
fn download_repo_file(repo_id: &str, file: &str) -> Result<std::path::PathBuf> {
|
||||||
use hf_hub::api::sync::ApiBuilder;
|
use hf_hub::api::sync::ApiBuilder;
|
||||||
let api = ApiBuilder::from_env()
|
let api = ApiBuilder::from_env()
|
||||||
@@ -69,43 +68,102 @@ fn download_repo_file(repo_id: &str, file: &str) -> Result<std::path::PathBuf> {
|
|||||||
.with_context(|| format!("download {file} for HuggingFace repo {repo_id:?}"))
|
.with_context(|| format!("download {file} for HuggingFace repo {repo_id:?}"))
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Load the `tokenizer_config.json` co-located with the tokenizer named by
|
/// List the files an HF repo ships. `None` (with a warning) when the listing
|
||||||
/// `source` (the same value passed to [`load`]). For a local
|
/// fails, e.g. offline with a warm cache; the caller then attempts each
|
||||||
/// `.../tokenizer.json` path this is the sibling file; for an HF repo id it is
|
/// download individually, which is cache-first.
|
||||||
/// downloaded from the same repo.
|
fn list_repo_files(repo_id: &str) -> Option<std::collections::HashSet<String>> {
|
||||||
///
|
use hf_hub::api::sync::ApiBuilder;
|
||||||
/// Returns `Ok(None)` when the model ships no `tokenizer_config.json` (rare but
|
let listing = ApiBuilder::from_env()
|
||||||
/// valid) — the caller then has no chat template and routes via raw prompt text.
|
.build()
|
||||||
pub fn load_tokenizer_config(source: &str) -> Result<Option<serde_json::Value>> {
|
.context("initialize HuggingFace Hub client")
|
||||||
let path = if Path::new(source).is_file() || looks_like_path(source) {
|
.and_then(|api| {
|
||||||
match Path::new(source).parent() {
|
api.model(repo_id.to_string())
|
||||||
Some(dir) => dir.join("tokenizer_config.json"),
|
.info()
|
||||||
None => return Ok(None),
|
.context("list repo files")
|
||||||
|
});
|
||||||
|
match listing {
|
||||||
|
Ok(info) => Some(info.siblings.into_iter().map(|s| s.rfilename).collect()),
|
||||||
|
Err(e) => {
|
||||||
|
tracing::warn!(repo = %repo_id, error = %format!("{e:#}"),
|
||||||
|
"could not list HuggingFace repo files; trying sibling downloads individually");
|
||||||
|
None
|
||||||
}
|
}
|
||||||
} else {
|
}
|
||||||
// HF repo id. The download error type doesn't distinguish a genuine
|
}
|
||||||
// 404 (repo ships no tokenizer_config.json — benign) from auth/network
|
|
||||||
// failures (wrong/expired HF_TOKEN, gated repo, timeout), so warn with
|
/// Files co-located with the tokenizer named by `source` (the same value passed
|
||||||
// the cause rather than asserting the benign case at debug: a swallowed
|
/// to [`load`]): siblings of a local `tokenizer.json`, or files of the same HF
|
||||||
// auth error here silently disables chat-template routing.
|
/// repo. The repo is listed once so only files it ships are downloaded; a file
|
||||||
match download_repo_file(source, "tokenizer_config.json") {
|
/// the model lacks resolves to `None` without a network round-trip.
|
||||||
|
pub struct ModelFiles {
|
||||||
|
source: String,
|
||||||
|
/// Directory of a local `tokenizer.json`; `None` for an HF repo id.
|
||||||
|
local_dir: Option<std::path::PathBuf>,
|
||||||
|
/// Repo listing; `None` when it failed and downloads are attempted blindly.
|
||||||
|
repo_files: Option<std::collections::HashSet<String>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl ModelFiles {
|
||||||
|
pub fn open(source: &str) -> Self {
|
||||||
|
let local_dir = (Path::new(source).is_file() || looks_like_path(source)).then(|| {
|
||||||
|
Path::new(source)
|
||||||
|
.parent()
|
||||||
|
.map_or_else(Default::default, Path::to_path_buf)
|
||||||
|
});
|
||||||
|
let repo_files = match local_dir {
|
||||||
|
Some(_) => None,
|
||||||
|
None => list_repo_files(source),
|
||||||
|
};
|
||||||
|
Self {
|
||||||
|
source: source.to_owned(),
|
||||||
|
local_dir,
|
||||||
|
repo_files,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn path(&self, file: &str) -> Option<std::path::PathBuf> {
|
||||||
|
let path = match &self.local_dir {
|
||||||
|
Some(dir) => dir.join(file),
|
||||||
|
None => {
|
||||||
|
if self
|
||||||
|
.repo_files
|
||||||
|
.as_ref()
|
||||||
|
.is_some_and(|files| !files.contains(file))
|
||||||
|
{
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
match download_repo_file(&self.source, file) {
|
||||||
Ok(p) => p,
|
Ok(p) => p,
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
tracing::warn!(repo = %source, error = %e,
|
tracing::warn!(repo = %self.source, %file, error = %format!("{e:#}"),
|
||||||
"could not download tokenizer_config.json; chat-template routing disabled for this model \
|
"could not download; chat-formatter detection may be degraded for this \
|
||||||
(expected if the repo ships none — otherwise check HF_TOKEN / network for a gated or private repo)");
|
model (check HF_TOKEN / network for a gated or private repo)");
|
||||||
return Ok(None);
|
return None;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
if !path.is_file() {
|
path.is_file().then_some(path)
|
||||||
return Ok(None);
|
|
||||||
}
|
}
|
||||||
let bytes = std::fs::read(&path)
|
|
||||||
.with_context(|| format!("read tokenizer_config.json at {}", path.display()))?;
|
/// Read the text `file`; `None` when the model ships no such file.
|
||||||
let value = serde_json::from_slice(&bytes)
|
pub fn text(&self, file: &str) -> Result<Option<String>> {
|
||||||
.with_context(|| format!("parse tokenizer_config.json at {}", path.display()))?;
|
self.path(file)
|
||||||
Ok(Some(value))
|
.map(|p| std::fs::read_to_string(&p).with_context(|| format!("read {}", p.display())))
|
||||||
|
.transpose()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Parse the JSON `file`; `None` when the model ships no such file.
|
||||||
|
pub fn json(&self, file: &str) -> Result<Option<serde_json::Value>> {
|
||||||
|
self.text(file)?
|
||||||
|
.map(|text| serde_json::from_str(&text).with_context(|| format!("parse {file}")))
|
||||||
|
.transpose()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Load the sibling tokenizer config for the current chat formatter.
|
||||||
|
pub fn load_tokenizer_config(source: &str) -> Result<Option<serde_json::Value>> {
|
||||||
|
ModelFiles::open(source).json("tokenizer_config.json")
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn encode(t: &Tokenizer, text: &str) -> Result<Vec<u32>> {
|
pub fn encode(t: &Tokenizer, text: &str) -> Result<Vec<u32>> {
|
||||||
@@ -137,3 +195,42 @@ pub fn decode_complete(t: &Tokenizer, ids: &[u32], skip_special: bool) -> Result
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod model_files_tests {
|
||||||
|
use super::ModelFiles;
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn reads_sibling_json_and_template_files() {
|
||||||
|
let dir = tempfile::tempdir().unwrap();
|
||||||
|
let tokenizer = dir.path().join("tokenizer.json");
|
||||||
|
std::fs::write(&tokenizer, "{}").unwrap();
|
||||||
|
std::fs::write(dir.path().join("config.json"), r#"{"model_type":"llama"}"#).unwrap();
|
||||||
|
std::fs::write(dir.path().join("chat_template.jinja"), "{{ messages }}").unwrap();
|
||||||
|
|
||||||
|
let files = ModelFiles::open(tokenizer.to_str().unwrap());
|
||||||
|
assert_eq!(
|
||||||
|
files.json("config.json").unwrap(),
|
||||||
|
Some(json!({"model_type":"llama"}))
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
files.text("chat_template.jinja").unwrap().as_deref(),
|
||||||
|
Some("{{ messages }}")
|
||||||
|
);
|
||||||
|
assert!(files.json("tokenizer_config.json").unwrap().is_none());
|
||||||
|
assert!(files.text("missing.jinja").unwrap().is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn invalid_json_is_an_error_instead_of_a_missing_file() {
|
||||||
|
let dir = tempfile::tempdir().unwrap();
|
||||||
|
let tokenizer = dir.path().join("tokenizer.json");
|
||||||
|
std::fs::write(&tokenizer, "{}").unwrap();
|
||||||
|
std::fs::write(dir.path().join("config.json"), "invalid JSON").unwrap();
|
||||||
|
|
||||||
|
let files = ModelFiles::open(tokenizer.to_str().unwrap());
|
||||||
|
let error = files.json("config.json").unwrap_err();
|
||||||
|
assert!(error.to_string().contains("config.json"));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user