use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet};
use std::io::{IsTerminal, Write as _};
use std::path::{Path, PathBuf};
use std::process::ExitCode;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use clap::{
ArgAction, ArgGroup, Args, CommandFactory, FromArgMatches, Parser, Subcommand, ValueEnum,
};
use tracing_subscriber::EnvFilter;
use hf_fetch_model::cache;
use hf_fetch_model::cache_layout;
use hf_fetch_model::discover;
use hf_fetch_model::header_cache;
use hf_fetch_model::inspect;
use hf_fetch_model::peek;
use hf_fetch_model::progress::IndicatifProgress;
use hf_fetch_model::repo;
use hf_fetch_model::{
DownloadPlan, FetchConfig, FetchConfigBuilder, FetchError, Filter, HttpRangeReader, RangeStats,
compile_glob_patterns, file_matches, has_glob_chars,
};
#[path = "../format.rs"]
mod format;
#[path = "../gpu_check.rs"]
mod gpu_check;
use format::format_size;
#[must_use]
fn apply_timeout_overrides(
builder: FetchConfigBuilder,
per_file_secs: Option<u64>,
total_secs: Option<u64>,
) -> FetchConfigBuilder {
let mut builder = builder;
if let Some(secs) = per_file_secs {
builder = builder.timeout_per_file(Duration::from_secs(secs));
}
if let Some(secs) = total_secs {
builder = builder.timeout_total(Duration::from_secs(secs));
}
builder
}
#[derive(Parser)]
#[command(
name = "hf-fetch-model",
bin_name = "hf-fm",
version,
about,
before_help = concat!("hf-fetch-model v", env!("CARGO_PKG_VERSION"))
)]
#[command(args_conflicts_with_subcommands = true)]
struct Cli {
#[command(subcommand)]
command: Option<Commands>,
#[command(flatten)]
download: DownloadArgs,
}
#[derive(Args)]
struct DownloadArgs {
#[arg(short, long)]
verbose: bool,
#[arg(value_name = "REPO_ID")]
repo_id: Option<String>,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long, action = clap::ArgAction::Append)]
filter: Vec<String>,
#[arg(long, action = clap::ArgAction::Append)]
exclude: Vec<String>,
#[arg(long, value_enum)]
preset: Option<Preset>,
#[arg(long)]
output_dir: Option<PathBuf>,
#[arg(long)]
concurrency: Option<usize>,
#[arg(long)]
chunk_threshold_mib: Option<u64>,
#[arg(long)]
connections_per_file: Option<usize>,
#[arg(long)]
timeout_per_file_secs: Option<u64>,
#[arg(long)]
timeout_total_secs: Option<u64>,
#[arg(long)]
dry_run: bool,
#[arg(long)]
flat: bool,
}
#[derive(Subcommand)]
enum Commands {
#[command(after_help = "Examples:
hf-fm list-families # default grouped view
hf-fm list-families --show quant # add a quant column
hf-fm list-families --tag bitsandbytes # filter to bitsandbytes-tagged repos
hf-fm list-families --show quant --tag gguf # combine
See also: hf-fm discover, hf-fm du")]
ListFamilies {
#[arg(long, value_delimiter = ',', value_enum)]
show: Vec<ShowFamiliesColumn>,
#[arg(long)]
tag: Option<String>,
},
#[command(after_help = "Examples:
hf-fm discover --limit 100 # top families by download count
hf-fm discover --tag bitsandbytes # only models carrying this tag
hf-fm discover --tag gguf --limit 200 # tag composes with --limit
See also: hf-fm list-families, hf-fm search")]
Discover {
#[arg(long, default_value = "500")]
limit: usize,
#[arg(long)]
tag: Option<String>,
},
#[command(after_help = "Examples:
hf-fm search \"fp4\" --tag bitsandbytes # text-match AND tag-match
hf-fm search \"llama\" --tag gguf --limit 5 # tag composes with other filters
hf-fm search \"qwen,3B\" --exact # exact repo-id match
hf-fm search \"fp4\" --show tags,size # enrich rows with tag list + total size
See also: hf-fm list-families, hf-fm discover")]
Search {
query: String,
#[arg(long, default_value = "20")]
limit: usize,
#[arg(long)]
exact: bool,
#[arg(long)]
library: Option<String>,
#[arg(long)]
pipeline: Option<String>,
#[arg(long)]
tag: Option<String>,
#[arg(long, value_delimiter = ',', value_enum)]
show: Vec<ShowColumn>,
#[arg(long)]
token: Option<String>,
},
#[command(after_help = "Examples:\n \
hf-fm quants poolside/Laguna-XS-2.1 # sorted quant table\n \
hf-fm quants poolside/Laguna-XS-2.1 --fits 16GiB --reserve 2.5GiB # offload-aware fit plan\n \
hf-fm quants poolside/Laguna-XS-2.1 --json # for scripting\n\n\
--fits skips header inspection for candidates that already fit under\n\
budget without offload; only over-budget MoE GGUF files are\n\
inspected (against the internal blk.*.*_exps.weight pattern) to\n\
compute a --n-cpu-moe plan.\n\n\
See also: hf-fm inspect <repo> <file> --group-by <PATTERN> # the rollup --fits uses internally")]
Quants {
repo_id: String,
#[arg(long, value_name = "SIZE", value_parser = parse_size_arg)]
fits: Option<u64>,
#[arg(long, value_name = "SIZE", value_parser = parse_size_arg, requires = "fits")]
reserve: Option<u64>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
json: bool,
},
Info {
repo_id: String,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
json: bool,
#[arg(long, default_value = "40")]
lines: usize,
},
DownloadFile {
#[arg(short, long)]
verbose: bool,
repo_id: String,
filename: String,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
output_dir: Option<PathBuf>,
#[arg(long)]
chunk_threshold_mib: Option<u64>,
#[arg(long)]
connections_per_file: Option<usize>,
#[arg(long)]
timeout_per_file_secs: Option<u64>,
#[arg(long)]
timeout_total_secs: Option<u64>,
#[arg(long)]
dry_run: bool,
#[arg(long)]
flat: bool,
},
Status {
repo_id: Option<String>,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long, value_enum)]
preset: Option<Preset>,
#[arg(long)]
json: bool,
},
Diff {
repo_a: String,
repo_b: String,
#[arg(long)]
revision_a: Option<String>,
#[arg(long)]
revision_b: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
cached: bool,
#[arg(long)]
filter: Option<String>,
#[arg(long, conflicts_with_all = ["dtypes", "collapse"])]
summary: bool,
#[arg(long, conflicts_with_all = ["summary", "collapse"])]
dtypes: bool,
#[arg(long, conflicts_with_all = ["summary", "dtypes"])]
collapse: bool,
#[arg(long)]
limit: Option<usize>,
#[arg(long)]
json: bool,
},
DiffConfig {
repo_a: String,
repo_b: String,
#[arg(long)]
revision_a: Option<String>,
#[arg(long)]
revision_b: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
cached: bool,
#[arg(long)]
all: bool,
#[arg(long)]
json: bool,
},
Du {
repo_id: Option<String>,
#[arg(long)]
age: bool,
#[arg(long, conflicts_with = "repo_id")]
tree: bool,
#[arg(long)]
json: bool,
},
#[command(after_help = "Examples:\n \
hf-fm inspect <repo> # inspect every .safetensors in the repo\n \
hf-fm inspect <repo> --filter blocks.0. # matched tensor names (per shard/file)\n \
hf-fm inspect <repo> --list # list tensor files (no headers read)\n \
hf-fm inspect <repo> 3 # inspect file #3 from --list\n \
hf-fm inspect <repo> --pick # pick the file interactively\n \
hf-fm inspect <repo> fluxV13 --pick --dtypes # substring narrows, then pick\n \
hf-fm inspect <repo> model.safetensors --tree # hierarchical view of one file\n \
hf-fm inspect <repo> model.gguf --group-by 'blk.*.ffn_*_exps' # MoE expert-byte rollup\n \
hf-fm inspect <repo> --check-gpu # GPU-fit verdict for the whole repo\n \
hf-fm inspect <repo> model.gguf # remote GGUF, no download (v0.11.2+)\n \
hf-fm inspect <repo> layer_9/width_16k/.../params.npz # remote NPZ, no download (v0.11.0+)\n \
hf-fm inspect <repo> pytorch_model.pth # remote PTH, no download (v0.11.4+)\n\n\
Indices returned by --list are stable as long as the repo has not\n\
changed remotely between invocations. Pass --revision <sha> on both\n\
--list and the follow-up run to lock the view end-to-end.\n\n\
--check-gpu reads device 0's VRAM via hypomnesis (NVML / DXGI) and\n\
reports a one-line fit verdict against the model's weight bytes.\n\
Pass --check-gpu N to target a specific device.\n\n\
For a walkthrough on a real 4-shard model, see\n\
docs/tutorials/inspect-before-downloading.md\n\n\
See also: hf-fm diff <A> <B> # compare two repos' tensor layouts")]
Inspect {
repo_id: String,
filename: Option<String>,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long)]
cached: bool,
#[arg(long, conflicts_with = "cached")]
cache_headers: bool,
#[arg(long, conflicts_with_all = ["filename", "no_metadata", "json", "filter", "dtypes", "limit", "tree", "group_by"])]
list: bool,
#[arg(long, conflicts_with = "list")]
pick: bool,
#[arg(long)]
no_metadata: bool,
#[arg(long)]
json: bool,
#[arg(long)]
filter: Option<String>,
#[arg(long)]
dtypes: bool,
#[arg(long, value_name = "PATTERN", conflicts_with = "dtypes")]
group_by: Option<String>,
#[arg(long)]
limit: Option<usize>,
#[arg(long, conflicts_with_all = ["dtypes", "limit", "group_by"])]
tree: bool,
#[arg(
long,
num_args = 0..=1,
default_missing_value = "0",
value_name = "N",
conflicts_with = "list"
)]
check_gpu: Option<u32>,
#[arg(long, value_name = "N", requires = "check_gpu")]
context: Option<u32>,
},
Peek {
repo_id: String,
filename: String,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long, value_name = "N", conflicts_with = "tail")]
head: Option<u64>,
#[arg(long, value_name = "N", conflicts_with_all = ["head", "gunzip"])]
tail: Option<u64>,
#[arg(long)]
bytes: bool,
#[arg(long, conflicts_with_all = ["no_gunzip", "tail"])]
gunzip: bool,
#[arg(long, conflicts_with = "gunzip")]
no_gunzip: bool,
#[arg(long, value_name = "SIZE", value_parser = parse_size_arg, default_value = "10MiB")]
max: u64,
},
ListFiles {
repo_id: String,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
#[arg(long, action = clap::ArgAction::Append)]
filter: Vec<String>,
#[arg(long, action = clap::ArgAction::Append)]
exclude: Vec<String>,
#[arg(long, value_enum)]
preset: Option<Preset>,
#[arg(long)]
no_checksum: bool,
#[arg(long)]
show_cached: bool,
#[arg(long)]
json: bool,
},
Cache {
#[command(subcommand)]
subcommand: CacheCommands,
},
}
#[derive(Subcommand)]
enum CacheCommands {
CleanPartial {
repo_id: Option<String>,
#[arg(long)]
yes: bool,
#[arg(long)]
dry_run: bool,
},
Delete {
repo_id: String,
#[arg(long)]
yes: bool,
},
#[command(group(
ArgGroup::new("gc_criteria")
.args(["older_than", "max_size"])
.required(true)
.multiple(true)
))]
Gc {
#[arg(long, value_name = "DAYS")]
older_than: Option<u64>,
#[arg(long, value_name = "SIZE", value_parser = parse_size_arg)]
max_size: Option<u64>,
#[arg(long = "except", value_name = "REPO_ID", action = ArgAction::Append)]
except: Vec<String>,
#[arg(long)]
dry_run: bool,
#[arg(long)]
yes: bool,
#[arg(long)]
list_kept: bool,
},
Path {
repo_id: String,
#[arg(long)]
revision: Option<String>,
},
Verify {
repo_id: String,
#[arg(long)]
revision: Option<String>,
#[arg(long)]
token: Option<String>,
},
}
#[derive(Clone, ValueEnum)]
enum Preset {
Safetensors,
Gguf,
Npz,
Pth,
ConfigOnly,
}
#[derive(Clone, Copy, PartialEq, Eq, ValueEnum)]
enum ShowColumn {
Tags,
Size,
}
#[derive(Clone, Copy, PartialEq, Eq, ValueEnum)]
enum ShowFamiliesColumn {
Quant,
}
const fn preset_name(preset: &Preset) -> &'static str {
match preset {
Preset::Safetensors => "safetensors",
Preset::Gguf => "gguf",
Preset::Npz => "npz",
Preset::Pth => "pth",
Preset::ConfigOnly => "config-only",
}
}
#[must_use]
fn sort_subcommands_alphabetically(mut cmd: clap::Command) -> clap::Command {
let mut names: Vec<String> = cmd
.get_subcommands()
.map(|sc| sc.get_name().to_owned()) .collect();
names.sort();
for (i, name) in names.iter().enumerate() {
cmd = cmd.mut_subcommand(name, |sc| {
sort_subcommands_alphabetically(sc).display_order(i)
});
}
cmd
}
fn main() -> ExitCode {
let cmd = sort_subcommands_alphabetically(Cli::command());
let matches = cmd.get_matches();
let cli = match Cli::from_arg_matches(&matches) {
Ok(cli) => cli,
Err(e) => e.exit(),
};
let verbose = match &cli.command {
Some(Commands::DownloadFile { verbose, .. }) => *verbose,
None => cli.download.verbose,
Some(
Commands::ListFamilies { .. }
| Commands::Discover { .. }
| Commands::Search { .. }
| Commands::Quants { .. }
| Commands::Info { .. }
| Commands::Status { .. }
| Commands::Diff { .. }
| Commands::DiffConfig { .. }
| Commands::Du { .. }
| Commands::Inspect { .. }
| Commands::Peek { .. }
| Commands::ListFiles { .. }
| Commands::Cache { .. },
) => false,
};
if verbose {
let filter = EnvFilter::try_from_default_env()
.unwrap_or_else(|_| EnvFilter::new("hf_fetch_model=debug"));
tracing_subscriber::fmt()
.with_env_filter(filter)
.with_target(false)
.with_writer(std::io::stderr)
.init();
}
match run(cli) {
Ok(()) => ExitCode::SUCCESS,
Err(FetchError::PartialDownload { path, failures }) => {
eprintln!();
eprintln!(
"error: {} {} failed to download:",
failures.len(),
pluralize(failures.len(), "file", "files")
);
for f in &failures {
eprintln!(" - {}: {}", f.filename, f.reason);
}
if let Some(p) = path {
eprintln!();
eprintln!("Partial download at: {}", p.display());
}
let any_retryable = failures.iter().any(|f| f.retryable);
if any_retryable {
eprintln!();
eprintln!(
"hint: re-run the same command to retry failed files \
(already-downloaded files will be skipped)"
);
}
ExitCode::FAILURE
}
Err(FetchError::RepoNotFound { ref repo_id }) => {
eprintln!(
"error: {e}",
e = FetchError::RepoNotFound {
repo_id: repo_id.clone()
}
);
let search_term = repo_id.split('/').nth(1).unwrap_or(repo_id.as_str());
eprintln!("hint: try `hf-fm search {search_term}` to find matching models");
ExitCode::FAILURE
}
Err(e) => {
eprintln!("error: {e}");
ExitCode::FAILURE
}
}
}
#[allow(clippy::too_many_lines)]
fn run(cli: Cli) -> Result<(), FetchError> {
match cli.command {
Some(Commands::ListFamilies { show, tag }) => run_list_families(&show, tag.as_deref()),
Some(Commands::Discover { limit, tag }) => run_discover(limit, tag.as_deref()),
Some(Commands::Search {
query,
limit,
exact,
library,
pipeline,
tag,
show,
token,
}) => run_search(
query.as_str(),
limit,
exact,
library.as_deref(),
pipeline.as_deref(),
tag.as_deref(),
&show,
token.as_deref(),
),
Some(Commands::Quants {
repo_id,
fits,
reserve,
token,
json,
}) => run_quants(repo_id.as_str(), fits, reserve, token.as_deref(), json),
Some(Commands::Info {
repo_id,
revision,
token,
json,
lines,
}) => run_info(
repo_id.as_str(),
revision.as_deref(),
token.as_deref(),
json,
lines,
),
Some(Commands::DownloadFile {
verbose: _,
repo_id,
filename,
revision,
token,
output_dir,
chunk_threshold_mib,
connections_per_file,
timeout_per_file_secs,
timeout_total_secs,
dry_run,
flat,
}) => run_download_file(DownloadFileParams {
repo_id: repo_id.as_str(),
filename: filename.as_str(),
revision: revision.as_deref(),
token: token.as_deref(),
output_dir,
chunk_threshold_mib,
connections_per_file,
timeout_per_file_secs,
timeout_total_secs,
dry_run,
flat,
}),
Some(Commands::Status {
repo_id: Some(repo_id),
revision,
token,
preset,
json,
}) => run_status(
repo_id.as_str(),
revision.as_deref(),
token.as_deref(),
preset.as_ref(),
json,
),
Some(Commands::Status {
repo_id: None,
json,
..
}) => run_status_all(json),
Some(Commands::Diff {
repo_a,
repo_b,
revision_a,
revision_b,
token,
cached,
filter,
summary,
dtypes,
collapse,
limit,
json,
}) => run_diff(
repo_a.as_str(),
repo_b.as_str(),
revision_a.as_deref(),
revision_b.as_deref(),
token.as_deref(),
cached,
filter.as_deref(),
summary,
dtypes,
collapse,
limit,
json,
),
Some(Commands::DiffConfig {
repo_a,
repo_b,
revision_a,
revision_b,
token,
cached,
all,
json,
}) => run_diff_config(
repo_a.as_str(),
repo_b.as_str(),
revision_a.as_deref(),
revision_b.as_deref(),
token.as_deref(),
cached,
all,
json,
),
Some(Commands::Du {
repo_id: Some(repo_id),
age: _,
tree: _, json,
}) => {
let resolved = resolve_du_arg(repo_id.as_str())?;
run_du_repo(resolved.as_str(), json)
}
Some(Commands::Du {
repo_id: None,
age,
tree: true,
json,
}) => run_du_tree(age, json),
Some(Commands::Du {
repo_id: None,
age,
tree: false,
json,
}) => run_du(age, json),
Some(Commands::Inspect {
repo_id,
filename,
revision,
token,
cached,
cache_headers,
list,
pick,
no_metadata,
json,
filter,
dtypes,
group_by,
limit,
tree,
check_gpu,
context,
}) => run_inspect(
repo_id.as_str(),
filename.as_deref(),
revision.as_deref(),
token.as_deref(),
cached,
cache_headers,
list,
pick,
no_metadata,
json,
filter.as_deref(),
dtypes,
group_by.as_deref(),
limit,
tree,
check_gpu,
context,
),
Some(Commands::Peek {
repo_id,
filename,
revision,
token,
head,
tail,
bytes,
gunzip,
no_gunzip,
max,
}) => run_peek(
repo_id.as_str(),
filename.as_str(),
revision.as_deref(),
token.as_deref(),
head,
tail,
bytes,
gunzip,
no_gunzip,
max,
),
Some(Commands::ListFiles {
repo_id,
revision,
token,
filter,
exclude,
preset,
no_checksum,
show_cached,
json,
}) => run_list_files(
repo_id.as_str(),
revision.as_deref(),
token.as_deref(),
&filter,
&exclude,
preset.as_ref(),
no_checksum,
show_cached,
json,
),
Some(Commands::Cache { subcommand }) => match subcommand {
CacheCommands::CleanPartial {
repo_id,
yes,
dry_run,
} => {
let resolved = repo_id.map(|r| resolve_du_arg(r.as_str())).transpose()?;
run_cache_clean_partial(resolved.as_deref(), yes, dry_run)
}
CacheCommands::Delete { repo_id, yes } => {
let resolved = resolve_du_arg(repo_id.as_str())?;
run_cache_delete(resolved.as_str(), yes)
}
CacheCommands::Gc {
older_than,
max_size,
except,
dry_run,
yes,
list_kept,
} => run_cache_gc(older_than, max_size, except, dry_run, yes, list_kept),
CacheCommands::Path { repo_id, revision } => {
let resolved = resolve_du_arg(repo_id.as_str())?;
run_cache_path(resolved.as_str(), revision.as_deref())
}
CacheCommands::Verify {
repo_id,
revision,
token,
} => {
let resolved = resolve_du_arg(repo_id.as_str())?;
run_cache_verify(resolved.as_str(), revision.as_deref(), token.as_deref())
}
},
None => run_download(cli.download),
}
}
struct NonTtyProgress {
last_report: Mutex<Instant>,
last_bucket: Mutex<HashMap<String, u64>>,
}
impl NonTtyProgress {
fn new() -> Self {
Self {
last_report: Mutex::new(Instant::now()),
last_bucket: Mutex::new(HashMap::new()),
}
}
fn handle(&self, event: &hf_fetch_model::progress::ProgressEvent) {
if event.percent >= 100.0 {
return;
}
#[allow(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::as_conversions
)]
let bucket = (event.percent / 10.0) as u64;
let elapsed_ok = self
.last_report
.lock()
.is_ok_and(|guard| guard.elapsed().as_secs() >= 5);
let bucket_crossed = self.last_bucket.lock().is_ok_and(|mut map| {
let prev = map.entry(event.filename.clone()).or_insert(0);
if bucket > *prev {
*prev = bucket;
true
} else {
false
}
});
if elapsed_ok || bucket_crossed {
if let Ok(mut ts) = self.last_report.lock() {
*ts = Instant::now();
}
#[allow(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::as_conversions
)]
let pct = event.percent as u64;
eprintln!(
"[hf-fm] {}: {}/{} ({pct}%)",
event.filename,
format_size(event.bytes_downloaded),
format_size(event.bytes_total)
);
}
}
}
#[allow(clippy::too_many_lines)]
fn run_download(args: DownloadArgs) -> Result<(), FetchError> {
let dry_run = args.dry_run;
let repo_id = args.repo_id.as_deref().ok_or_else(|| {
FetchError::InvalidArgument(
"REPO_ID is required for download. Usage: hf-fm <REPO_ID>".to_owned(),
)
})?;
if !repo_id.contains('/') {
return Err(FetchError::InvalidArgument(format!(
"invalid REPO_ID \"{repo_id}\": expected \"org/model\" format (e.g., \"EleutherAI/pythia-1.4b\")"
)));
}
if dry_run {
return run_dry_run(repo_id, &args);
}
let repo_id = repo_id.to_owned();
let flat = args.flat;
let flat_target = if flat { args.output_dir.clone() } else { None };
let mut builder = match args.preset {
Some(Preset::Safetensors) => Filter::safetensors(),
Some(Preset::Gguf) => Filter::gguf(),
Some(Preset::Npz) => Filter::npz(),
Some(Preset::Pth) => Filter::pth(),
Some(Preset::ConfigOnly) => Filter::config_only(),
None => FetchConfig::builder(),
};
if let Some(ref preset) = args.preset {
warn_redundant_filters(preset, &args.filter);
}
if let Some(rev) = args.revision.as_deref() {
builder = builder.revision(rev);
}
if let Some(tok) = args.token.as_deref() {
builder = builder.token(tok);
} else {
builder = builder.token_from_env();
}
for pattern in &args.filter {
builder = builder.filter(pattern.as_str());
}
for pattern in &args.exclude {
builder = builder.exclude(pattern.as_str());
}
if let Some(c) = args.concurrency {
builder = builder.concurrency(c);
}
if let Some(ct) = args.chunk_threshold_mib {
builder = builder.chunk_threshold(ct.saturating_mul(1024 * 1024));
}
if let Some(cpf) = args.connections_per_file {
builder = builder.connections_per_file(cpf);
}
builder = apply_timeout_overrides(builder, args.timeout_per_file_secs, args.timeout_total_secs);
if !flat && let Some(dir) = args.output_dir {
builder = builder.output_dir(dir);
}
let is_tty = std::io::stderr().is_terminal();
let indicatif = if is_tty {
let p = Arc::new(IndicatifProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
Some(p)
} else {
let p = Arc::new(NonTtyProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
None
};
let config = builder.build()?;
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let start = Instant::now();
if flat {
let repo_id_for_snapshot = repo_id.clone();
let outcome = rt.block_on(hf_fetch_model::download_files_with_config(repo_id, &config))?;
let elapsed = start.elapsed();
if let Some(ref p) = indicatif {
p.finish();
}
if let Err(e) = write_download_snapshot(
repo_id_for_snapshot.as_str(), args.preset.as_ref(),
&args.filter,
&args.exclude,
args.revision.as_deref(), ) {
eprintln!("warning: could not write snapshot sidecar: {e}");
}
let file_map = outcome.inner();
let target_dir = resolve_flat_target(flat_target.as_deref())?;
let flat_paths = flatten_files(file_map, &target_dir)?;
println!(
"{} {} copied to {}:",
flat_paths.len(),
pluralize(flat_paths.len(), "file", "files"),
target_dir.display()
);
for p in &flat_paths {
println!(" {}", p.display());
}
print_download_summary(&target_dir, elapsed);
} else {
let repo_id_for_snapshot = repo_id.clone();
let outcome = rt.block_on(hf_fetch_model::download_with_config(repo_id, &config))?;
let elapsed = start.elapsed();
if let Some(ref p) = indicatif {
p.finish();
}
if let Err(e) = write_download_snapshot(
repo_id_for_snapshot.as_str(), args.preset.as_ref(),
&args.filter,
&args.exclude,
args.revision.as_deref(), ) {
eprintln!("warning: could not write snapshot sidecar: {e}");
}
if outcome.is_cached() {
println!("Cached at: {}", outcome.inner().display());
} else {
println!("Downloaded to: {}", outcome.inner().display());
print_download_summary(outcome.inner(), elapsed);
}
}
Ok(())
}
fn write_download_snapshot(
repo_id: &str,
preset: Option<&Preset>,
filter: &[String],
exclude: &[String],
revision: Option<&str>,
) -> Result<(), FetchError> {
let cache_root = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_root, repo_id);
let snapshot = cache::Snapshot {
version: cache::SNAPSHOT_VERSION,
revision: revision.unwrap_or("main").to_owned(),
preset: preset.map(|p| preset_name(p).to_owned()),
filter: filter.to_vec(),
exclude: exclude.to_vec(),
};
cache::write_snapshot(&repo_dir, &snapshot)
}
fn run_dry_run(repo_id: &str, args: &DownloadArgs) -> Result<(), FetchError> {
let mut builder = match args.preset {
Some(Preset::Safetensors) => Filter::safetensors(),
Some(Preset::Gguf) => Filter::gguf(),
Some(Preset::Npz) => Filter::npz(),
Some(Preset::Pth) => Filter::pth(),
Some(Preset::ConfigOnly) => Filter::config_only(),
None => FetchConfig::builder(),
};
if let Some(ref preset) = args.preset {
warn_redundant_filters(preset, &args.filter);
}
if let Some(rev) = args.revision.as_deref() {
builder = builder.revision(rev);
}
if let Some(tok) = args.token.as_deref() {
builder = builder.token(tok);
} else {
builder = builder.token_from_env();
}
for pattern in &args.filter {
builder = builder.filter(pattern.as_str());
}
for pattern in &args.exclude {
builder = builder.exclude(pattern.as_str());
}
if let Some(ref dir) = args.output_dir {
builder = builder.output_dir(dir.clone());
}
let config = builder.build()?;
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let plan = rt.block_on(hf_fetch_model::download_plan(repo_id, &config))?;
println!(" Repo: {}", plan.repo_id);
println!(" Revision: {}", plan.revision);
if args.preset.is_some() || !args.filter.is_empty() {
println!(" Filter: active (preset or --filter)");
}
if args.flat {
let target = resolve_flat_target(args.output_dir.as_deref())?;
println!(
" Flat: {} (files will be copied here)",
target.display()
);
}
println!();
render_download_plan(&plan)
}
fn render_download_plan(plan: &DownloadPlan) -> Result<(), FetchError> {
let fw = plan
.files
.iter()
.map(|fp| fp.filename.len())
.max()
.unwrap_or(4)
.max(4); let row_width = fw + 2 + 10 + 2 + 11;
println!(" {:<fw$} {:>10} Status", "File", "Size");
println!(
" {:\u{2500}<fw$} {:\u{2500}<10} {:\u{2500}<12}",
"", "", ""
);
for fp in &plan.files {
let status = if fp.cached {
"cached \u{2713}"
} else {
"to download"
};
println!(
" {:<fw$} {:>10} {status}",
fp.filename,
format_size(fp.size)
);
}
println!("{:\u{2500}<row_width$}", " ");
let cached_count = plan.files.len() - plan.files_to_download();
let to_dl = plan.files_to_download();
println!(
" Total: {} ({} {}, {} cached, {} to download)",
format_size(plan.total_bytes),
plan.files.len(),
pluralize(plan.files.len(), "file", "files"),
cached_count,
to_dl
);
println!(" Download: {}", format_size(plan.download_bytes));
if !plan.fully_cached() {
let rec = plan.recommended_config()?;
println!();
println!(" Recommended config:");
println!(" concurrency: {}", rec.concurrency());
println!(" connections/file: {}", rec.connections_per_file());
if rec.chunk_threshold() == u64::MAX {
println!(" chunk threshold: disabled (single-connection per file)");
} else {
println!(
" chunk threshold: {} MiB",
rec.chunk_threshold() / 1_048_576
);
}
}
Ok(())
}
struct DownloadFileParams<'a> {
repo_id: &'a str,
filename: &'a str,
revision: Option<&'a str>,
token: Option<&'a str>,
output_dir: Option<PathBuf>,
chunk_threshold_mib: Option<u64>,
connections_per_file: Option<usize>,
timeout_per_file_secs: Option<u64>,
timeout_total_secs: Option<u64>,
dry_run: bool,
flat: bool,
}
fn run_download_file(params: DownloadFileParams<'_>) -> Result<(), FetchError> {
let DownloadFileParams {
repo_id,
filename,
revision,
token,
output_dir,
chunk_threshold_mib,
connections_per_file,
timeout_per_file_secs,
timeout_total_secs,
dry_run,
flat,
} = params;
if !repo_id.contains('/') {
return Err(FetchError::InvalidArgument(format!(
"invalid REPO_ID \"{repo_id}\": expected \"org/model\" format (e.g., \"mntss/clt-gemma-2-2b-426k\")"
)));
}
if dry_run {
return run_download_file_dry_run(repo_id, filename, revision, token, output_dir, flat);
}
if has_glob_chars(filename) {
return run_download_file_glob(DownloadFileParams {
repo_id,
filename,
revision,
token,
output_dir,
chunk_threshold_mib,
connections_per_file,
timeout_per_file_secs,
timeout_total_secs,
dry_run,
flat,
});
}
let flat_target = if flat { output_dir.clone() } else { None };
let mut builder = FetchConfig::builder();
if let Some(rev) = revision {
builder = builder.revision(rev);
}
if let Some(tok) = token {
builder = builder.token(tok);
} else {
builder = builder.token_from_env();
}
if let Some(ct) = chunk_threshold_mib {
builder = builder.chunk_threshold(ct.saturating_mul(1024 * 1024));
}
if let Some(cpf) = connections_per_file {
builder = builder.connections_per_file(cpf);
}
builder = apply_timeout_overrides(builder, timeout_per_file_secs, timeout_total_secs);
if !flat && let Some(dir) = output_dir {
builder = builder.output_dir(dir);
}
let is_tty = std::io::stderr().is_terminal();
let indicatif = if is_tty {
let p = Arc::new(IndicatifProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
Some(p)
} else {
let p = Arc::new(NonTtyProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
None
};
let config = builder.build()?;
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let start = Instant::now();
let outcome = rt.block_on(hf_fetch_model::download_file(
repo_id.to_owned(),
filename,
&config,
))?;
let elapsed = start.elapsed();
if let Some(ref p) = indicatif {
p.finish();
}
if flat {
let target_dir = resolve_flat_target(flat_target.as_deref())?;
let flat_path = flatten_single_file(outcome.inner(), &target_dir)?;
println!("Copied to: {}", flat_path.display());
} else if outcome.is_cached() {
println!("Cached at: {}", outcome.inner().display());
} else {
println!("Downloaded to: {}", outcome.inner().display());
print_download_summary(outcome.inner(), elapsed);
}
Ok(())
}
fn run_download_file_dry_run(
repo_id: &str,
filename: &str,
revision: Option<&str>,
token: Option<&str>,
output_dir: Option<PathBuf>,
flat: bool,
) -> Result<(), FetchError> {
let mut builder = FetchConfig::builder().filter(filename);
if let Some(rev) = revision {
builder = builder.revision(rev);
}
if let Some(tok) = token {
builder = builder.token(tok);
} else {
builder = builder.token_from_env();
}
let flat_target = if flat { output_dir.clone() } else { None };
if !flat && let Some(dir) = output_dir {
builder = builder.output_dir(dir);
}
let config = builder.build()?;
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let plan = rt.block_on(hf_fetch_model::download_plan(repo_id, &config))?;
if plan.files.is_empty() {
if has_glob_chars(filename) {
println!("No files matched pattern \"{filename}\" in {repo_id}");
return Ok(());
}
return Err(FetchError::InvalidArgument(format!(
"\"{filename}\" not found in {repo_id}"
)));
}
println!(" Repo: {}", plan.repo_id);
println!(" Revision: {}", plan.revision);
if flat {
let target = resolve_flat_target(flat_target.as_deref())?;
println!(
" Flat: {} (files will be copied here)",
target.display()
);
}
println!();
render_download_plan(&plan)
}
fn run_download_file_glob(params: DownloadFileParams<'_>) -> Result<(), FetchError> {
let DownloadFileParams {
repo_id,
filename: pattern,
revision,
token,
output_dir,
chunk_threshold_mib,
connections_per_file,
timeout_per_file_secs,
timeout_total_secs,
dry_run: _,
flat,
} = params;
let flat_target = if flat { output_dir.clone() } else { None };
let mut builder = FetchConfig::builder().filter(pattern);
if let Some(rev) = revision {
builder = builder.revision(rev);
}
if let Some(tok) = token {
builder = builder.token(tok);
} else {
builder = builder.token_from_env();
}
if let Some(ct) = chunk_threshold_mib {
builder = builder.chunk_threshold(ct.saturating_mul(1024 * 1024));
}
if let Some(cpf) = connections_per_file {
builder = builder.connections_per_file(cpf);
}
builder = apply_timeout_overrides(builder, timeout_per_file_secs, timeout_total_secs);
if !flat && let Some(dir) = output_dir {
builder = builder.output_dir(dir);
}
let is_tty = std::io::stderr().is_terminal();
let indicatif = if is_tty {
let p = Arc::new(IndicatifProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
Some(p)
} else {
let p = Arc::new(NonTtyProgress::new());
let handle = Arc::clone(&p);
builder = builder.on_progress(move |e| handle.handle(e));
None
};
let config = builder.build()?;
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let start = Instant::now();
let outcome = rt.block_on(hf_fetch_model::download_files_with_config(
repo_id.to_owned(),
&config,
))?;
let elapsed = start.elapsed();
if let Some(ref p) = indicatif {
p.finish();
}
let file_map = outcome.inner();
if file_map.is_empty() {
println!("No files matched pattern \"{pattern}\" in {repo_id}");
return Ok(());
}
if flat {
let target_dir = resolve_flat_target(flat_target.as_deref())?;
let flat_paths = flatten_files(file_map, &target_dir)?;
println!(
"{} {} copied to {}:",
flat_paths.len(),
pluralize(flat_paths.len(), "file", "files"),
target_dir.display()
);
for p in &flat_paths {
println!(" {}", p.display());
}
} else {
println!(
"{} {} matched pattern \"{pattern}\":",
file_map.len(),
pluralize(file_map.len(), "file", "files")
);
for (name, path) in file_map {
println!(" {name}: {}", path.display());
}
}
let elapsed_secs = elapsed.as_secs_f64();
if elapsed_secs > 0.0 {
println!(" completed in {elapsed_secs:.1}s");
}
Ok(())
}
#[allow(clippy::too_many_lines)]
fn run_list_families(show: &[ShowFamiliesColumn], tag: Option<&str>) -> Result<(), FetchError> {
let show_quant = show.contains(&ShowFamiliesColumn::Quant);
let cache_dir = cache::hf_cache_dir()?;
let mut families = cache::list_cached_families()?;
println!("Cache: {}", cache_dir.display());
println!();
if families.is_empty() {
println!("No model families found in local cache.");
return Ok(());
}
if let Some(tag_filter) = tag {
let lower_tag = tag_filter.to_lowercase();
let repo_ids: Vec<String> = families
.values()
.flat_map(|entries| entries.iter().map(|e| e.repo_id.clone()))
.collect();
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let tags_by_repo = rt.block_on(discover::fetch_tags_concurrent(repo_ids));
for entries in families.values_mut() {
entries.retain(|entry| {
let Some(tags) = tags_by_repo.get(entry.repo_id.as_str()) else {
return false;
};
tags.iter()
.any(|t| t.eq_ignore_ascii_case(lower_tag.as_str()))
});
}
families.retain(|_, entries| !entries.is_empty());
if families.is_empty() {
println!("No cached families match tag {tag_filter:?}.");
return Ok(());
}
}
let fw = families
.keys()
.map(String::len)
.max()
.unwrap_or(6)
.max(6) + 2;
let mw = families
.values()
.flat_map(|entries| entries.iter().map(|e| e.repo_id.len()))
.max()
.unwrap_or(6)
.max(6); let qw = if show_quant {
families
.values()
.flat_map(|entries| {
entries
.iter()
.map(|e| e.quant_method.as_deref().unwrap_or("\u{2014}").len())
})
.max()
.unwrap_or(5)
.max(5) + 2
} else {
0
};
if show_quant {
println!("{:<fw$}{:<qw$}Models", "Family", "Quant");
println!("{:-<fw$}{:-<qw$}{:-<mw$}", "", "", "");
} else {
println!("{:<fw$}Models", "Family");
println!("{:-<fw$}{:-<mw$}", "", "");
}
for (model_type, entries) in &families {
for (i, entry) in entries.iter().enumerate() {
let quant_cell = entry.quant_method.as_deref().unwrap_or("\u{2014}");
let family_cell = if i == 0 { model_type.as_str() } else { "" }; if show_quant {
println!("{family_cell:<fw$}{quant_cell:<qw$}{}", entry.repo_id);
} else {
println!("{family_cell:<fw$}{}", entry.repo_id);
}
}
}
Ok(())
}
fn run_discover(limit: usize, tag: Option<&str>) -> Result<(), FetchError> {
let families = cache::list_cached_families()?;
let local_types: HashSet<String> = families.into_keys().collect();
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let discovered = rt.block_on(discover::discover_new_families(&local_types, limit, tag))?;
if discovered.is_empty() {
match tag {
Some(t) => println!("No new model families found with tag {t:?}."),
None => println!("No new model families found."),
}
return Ok(());
}
match tag {
Some(t) => {
println!("New families with tag {t:?} not in local cache (top models by downloads):\n");
}
None => println!("New families not in local cache (top models by downloads):\n"),
}
let fw = discovered
.iter()
.map(|f| f.model_type.len())
.max()
.unwrap_or(6)
.max(6) + 2;
let mw = discovered
.iter()
.map(|f| f.top_model.len())
.max()
.unwrap_or(9)
.max(9); println!("{:<fw$}Top Model", "Family");
println!("{:-<fw$}{:-<mw$}", "", "");
for family in &discovered {
println!("{:<fw$}{}", family.model_type, family.top_model);
}
Ok(())
}
#[allow(clippy::too_many_arguments, clippy::too_many_lines)]
fn run_search(
query: &str,
limit: usize,
exact: bool,
library: Option<&str>,
pipeline: Option<&str>,
tag: Option<&str>,
show: &[ShowColumn],
token: Option<&str>,
) -> Result<(), FetchError> {
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let show_tags = show.contains(&ShowColumn::Tags);
let show_size = show.contains(&ShowColumn::Size);
let has_commas = query.contains(',');
let normalized = if has_commas {
query.replace('/', ",")
} else {
query.replace('/', " ")
};
let terms: Vec<&str> = normalized
.split(',')
.map(str::trim)
.filter(|t| !t.is_empty())
.collect();
let api_query = terms.first().copied().unwrap_or(normalized.as_str());
let filter_terms: Vec<String> = terms.iter().map(|t| t.to_lowercase()).collect();
let has_client_filter =
filter_terms.len() > 1 || library.is_some() || pipeline.is_some() || tag.is_some();
let api_limit = if has_client_filter {
limit.saturating_mul(5)
} else {
limit
};
let results = rt.block_on(discover::search_models(
api_query,
api_limit,
library,
pipeline,
tag,
token.as_deref(),
))?;
let has_multi_term = filter_terms.len() > 1;
let normalized_ids: Vec<String> = if has_multi_term {
results
.iter()
.map(|r| r.model_id.replace('/', " ").to_lowercase())
.collect()
} else {
Vec::new()
};
let filtered: Vec<&discover::SearchResult> = results
.iter()
.enumerate()
.filter(|(i, _)| {
if !has_multi_term {
return true;
}
#[allow(clippy::indexing_slicing)]
let id_normalized = &normalized_ids[*i];
filter_terms
.iter()
.all(|term| id_normalized.contains(term.as_str())) })
.map(|(_, r)| r)
.take(limit)
.collect();
let size_by_repo: HashMap<String, discover::RepoSizeSummary> =
if show_size && !filtered.is_empty() {
let repo_ids: Vec<String> = filtered.iter().map(|r| r.model_id.clone()).collect();
rt.block_on(discover::fetch_repo_size_summaries_concurrent(
repo_ids,
token.as_deref(),
))?
} else {
HashMap::new()
};
if exact {
let exact_match = filtered
.iter()
.find(|r| r.model_id.eq_ignore_ascii_case(query));
if let Some(matched) = exact_match {
println!("Exact match:\n");
let size_summary = size_by_repo.get(matched.model_id.as_str()).copied(); print_search_result(
matched,
matched.model_id.len(),
show_tags,
show_size,
size_summary,
);
match rt.block_on(discover::fetch_model_card(
matched.model_id.as_str(), )) {
Ok(card) => print_model_card(&card),
Err(e) => eprintln!("\n (could not fetch model card: {e})"),
}
println!("\n See also: hf-fm info {}", matched.model_id);
} else {
println!("No exact match for \"{query}\".");
if !filtered.is_empty() {
println!("\nDid you mean:\n");
let nw = filtered.iter().map(|r| r.model_id.len()).max().unwrap_or(0);
for result in &filtered {
let size_summary = size_by_repo.get(result.model_id.as_str()).copied(); print_search_result(result, nw, show_tags, show_size, size_summary);
}
}
}
} else {
if filtered.is_empty() {
println!("No models found matching \"{query}\".");
} else {
let nw = filtered.iter().map(|r| r.model_id.len()).max().unwrap_or(0);
println!("Models matching \"{query}\" (by downloads):\n");
for result in &filtered {
let size_summary = size_by_repo.get(result.model_id.as_str()).copied(); print_search_result(result, nw, show_tags, show_size, size_summary);
}
}
}
Ok(())
}
fn print_search_result(
result: &discover::SearchResult,
name_width: usize,
show_tags: bool,
show_size: bool,
size_summary: Option<discover::RepoSizeSummary>,
) {
let suffix = match (&result.library_name, &result.pipeline_tag) {
(Some(lib), Some(pipe)) => format!(" [{lib}, {pipe}]"),
(Some(lib), None) => format!(" [{lib}]"),
(None, Some(pipe)) => format!(" [{pipe}]"),
(None, None) => String::new(),
};
let downloads_label = if result.downloads == 1 {
"download"
} else {
"downloads"
};
let size_col = if show_size {
match size_summary {
Some(discover::RepoSizeSummary {
quant_alternatives: true,
size_min: Some(min),
size_max: Some(max),
..
}) => format!(" {} to {}", format_size(min), format_size(max)),
Some(summary) => format!(" {}", format_size(summary.total)),
None => " \u{2014}".to_owned(), }
} else {
String::new()
};
let tags_col = if show_tags && !result.tags.is_empty() {
format!(" tags: {}", result.tags.join(", "))
} else {
String::new()
};
println!(
" hf-fm {:<nw$} ({} {downloads_label}){suffix}{size_col}{tags_col}",
result.model_id,
format_downloads(result.downloads),
nw = name_width,
);
}
const MOE_EXPERT_PATTERN: &str = "blk.*.*_exps.weight";
const QUANT_SCHEME_BITS: &[(&str, f64)] = &[
("IQ1_S", 1.56),
("IQ1_M", 1.75),
("IQ2_XXS", 2.06),
("IQ2_XS", 2.31),
("IQ2_S", 2.5),
("IQ2_M", 2.7),
("Q2_K_S", 2.16),
("Q2_K", 2.63),
("IQ3_XXS", 3.06),
("IQ3_XS", 3.3),
("IQ3_S", 3.44),
("IQ3_M", 3.66),
("Q3_K_S", 3.44),
("Q3_K_M", 3.74),
("Q3_K_L", 4.03),
("Q3_K", 3.74),
("IQ4_XS", 4.25),
("IQ4_NL", 4.5),
("Q4_0", 4.5),
("Q4_1", 5.0),
("Q4_K_S", 4.58),
("Q4_K_M", 4.85),
("Q4_K", 4.85),
("Q5_0", 5.5),
("Q5_1", 6.0),
("Q5_K_S", 5.54),
("Q5_K_M", 5.68),
("Q5_K", 5.68),
("Q6_K", 6.56),
("Q8_0", 8.5),
("BF16", 16.0),
("F16", 16.0),
("F32", 32.0),
("NVFP4", 4.0),
("MXFP4", 4.0),
("FP8", 8.0),
("INT8", 8.0),
("INT4", 4.0),
];
fn bits_for_artifact(name: &str) -> Option<f64> {
let upper = name.to_ascii_uppercase();
QUANT_SCHEME_BITS
.iter()
.find(|(token, _)| upper.contains(token))
.map(|(_, bits)| *bits)
}
#[derive(Clone)]
struct QuantArtifactRow {
artifact: String,
size: u64,
repo: String,
is_gguf: bool,
bits: Option<f64>,
verification: discover::QuantVerification,
}
fn build_quant_rows(candidates: Vec<discover::QuantCandidate>) -> Vec<QuantArtifactRow> {
let mut rows = Vec::new();
for candidate in candidates {
let gguf_files: Vec<&repo::RepoFile> = candidate
.files
.iter()
.filter(|f| discover::is_gguf_filename(&f.filename))
.collect();
if gguf_files.is_empty() {
let total: u64 = candidate
.files
.iter()
.filter(|f| f.filename.to_ascii_lowercase().ends_with(".safetensors"))
.filter_map(|f| f.size)
.fold(0u64, u64::saturating_add);
if total == 0 {
continue; }
let short_name = candidate
.repo_id
.rsplit('/')
.next()
.unwrap_or(candidate.repo_id.as_str());
rows.push(QuantArtifactRow {
artifact: short_name.to_owned(),
size: total,
bits: bits_for_artifact(&candidate.repo_id),
is_gguf: false,
repo: candidate.repo_id.clone(),
verification: candidate.verification.clone(),
});
continue;
}
for f in &gguf_files {
rows.push(QuantArtifactRow {
artifact: f.filename.clone(),
size: f.size.unwrap_or(0),
bits: bits_for_artifact(&f.filename),
is_gguf: true,
repo: candidate.repo_id.clone(),
verification: candidate.verification.clone(),
});
}
}
rows.sort_by_key(|r| r.size);
rows
}
#[derive(Debug, Clone)]
enum FitVerdict {
FullGpu,
Offload {
n_cpu_moe: u32,
moved_bytes: u64,
resident_bytes: u64,
},
DoesNotFit { reason: String },
}
const FITS_INSPECT_CONCURRENCY: usize = 8;
fn trivial_fit_verdict(row: &QuantArtifactRow, budget: u64, reserve: u64) -> Option<FitVerdict> {
let budget_after_reserve = budget.saturating_sub(reserve);
if row.size <= budget_after_reserve {
return Some(FitVerdict::FullGpu);
}
if !row.is_gguf {
return Some(FitVerdict::DoesNotFit {
reason: "no offload mechanism for this format".to_owned(),
});
}
None
}
async fn compute_fit_verdicts_concurrent(
rows: &[QuantArtifactRow],
budget: u64,
reserve_bytes: u64,
token: Option<&str>,
) -> Vec<FitVerdict> {
let mut verdicts: Vec<Option<FitVerdict>> = Vec::with_capacity(rows.len());
let mut needs_inspection: Vec<(usize, QuantArtifactRow)> = Vec::new();
for (index, row) in rows.iter().enumerate() {
if let Some(verdict) = trivial_fit_verdict(row, budget, reserve_bytes) {
verdicts.push(Some(verdict));
} else {
verdicts.push(None);
needs_inspection.push((index, row.clone()));
}
}
let token_owned = token.map(String::from);
let results = discover::fan_out_bounded(
needs_inspection,
FITS_INSPECT_CONCURRENCY,
move |(index, row)| {
let token_owned = token_owned.clone();
async move {
let verdict =
compute_fit_verdict(&row, budget, reserve_bytes, token_owned.as_deref()).await;
Some((index, verdict))
}
},
)
.await;
for (index, verdict) in results.into_iter().flatten() {
if let Some(slot) = verdicts.get_mut(index) {
*slot = Some(verdict);
}
}
verdicts
.into_iter()
.map(|slot| {
slot.unwrap_or_else(|| FitVerdict::DoesNotFit {
reason: "offload check task did not run".to_owned(),
})
})
.collect()
}
async fn compute_fit_verdict(
row: &QuantArtifactRow,
budget: u64,
reserve: u64,
token: Option<&str>,
) -> FitVerdict {
if let Some(verdict) = trivial_fit_verdict(row, budget, reserve) {
return verdict;
}
let budget_after_reserve = budget.saturating_sub(reserve);
let (info, _source, _stats) =
match inspect::inspect_gguf(&row.repo, &row.artifact, token, None).await {
Ok(v) => v,
Err(e) => {
return FitVerdict::DoesNotFit {
reason: format!("offload plan unavailable: {e}"),
};
}
};
let Ok(matcher) = compile_group_by_pattern(MOE_EXPERT_PATTERN) else {
return FitVerdict::DoesNotFit {
reason: "internal offload pattern failed to compile".to_owned(),
};
};
let rollup = compute_group_by_rollup(&info.tensors, &matcher);
let (Some(layer_count), Some(per_layer)) = (rollup.layer_count, rollup.per_layer_bytes) else {
return FitVerdict::DoesNotFit {
reason: "no MoE experts to offload".to_owned(),
};
};
if per_layer == 0 {
return FitVerdict::DoesNotFit {
reason: "no MoE experts to offload".to_owned(),
};
}
compute_offload_plan(
row.size,
rollup.matched_bytes,
layer_count,
per_layer,
budget_after_reserve,
)
}
fn compute_offload_plan(
total_bytes: u64,
matched_bytes: u64,
layer_count: usize,
per_layer_bytes: u64,
budget_after_reserve: u64,
) -> FitVerdict {
let shortfall = total_bytes.saturating_sub(budget_after_reserve);
let n = shortfall.div_ceil(per_layer_bytes);
#[allow(clippy::as_conversions)]
let layer_count = layer_count as u64;
let n = n.min(layer_count);
let moved = n.saturating_mul(per_layer_bytes);
let resident = total_bytes.saturating_sub(moved);
if resident > budget_after_reserve {
FitVerdict::DoesNotFit {
reason: format!(
"does not fit even with full expert offload ({} non-expert weight)",
format_size(total_bytes.saturating_sub(matched_bytes))
),
}
} else {
#[allow(clippy::as_conversions, clippy::cast_possible_truncation)]
let n_cpu_moe = n as u32;
FitVerdict::Offload {
n_cpu_moe,
moved_bytes: moved,
resident_bytes: resident,
}
}
}
fn print_quants_table(rows: &[QuantArtifactRow], fits: Option<&[FitVerdict]>) {
let aw = rows
.iter()
.map(|r| r.artifact.len())
.max()
.unwrap_or(8)
.max(8);
println!();
if let Some(verdicts) = fits {
println!(
" {:<aw$} {:>10} {:>10} PLAN",
"ARTIFACT", "SIZE", "RESIDENT"
);
for (row, verdict) in rows.iter().zip(verdicts) {
let (resident, plan) = match verdict {
FitVerdict::FullGpu => (format_size(row.size), "full GPU".to_owned()),
FitVerdict::Offload {
n_cpu_moe,
moved_bytes,
resident_bytes,
} => (
format_size(*resident_bytes),
format!(
"--n-cpu-moe {n_cpu_moe} ({} -> RAM)",
format_size(*moved_bytes)
),
),
FitVerdict::DoesNotFit { reason } => {
("\u{2014}".to_owned(), format!("does not fit ({reason})"))
}
};
println!(
" {:<aw$} {:>10} {:>10} {plan}",
row.artifact,
format_size(row.size),
resident,
);
}
} else {
let rw = rows.iter().map(|r| r.repo.len()).max().unwrap_or(4).max(4); println!(
" {:<aw$} {:>10} {:<rw$} BITS",
"ARTIFACT", "SIZE", "REPO"
);
for row in rows {
let bits = row
.bits
.map_or_else(|| "?".to_owned(), |b| format!("~{b:.1}"));
println!(
" {:<aw$} {:>10} {:<rw$} {bits}",
row.artifact,
format_size(row.size),
row.repo,
);
}
}
}
#[derive(serde::Serialize)]
#[serde(rename_all = "snake_case")]
enum VerificationJson {
Verified,
Unverified,
CheckFailed,
}
#[derive(serde::Serialize)]
struct FitJson {
fits: bool,
resident_bytes: u64,
#[serde(skip_serializing_if = "Option::is_none")]
n_cpu_moe: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
moved_bytes: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
reason: Option<String>,
}
#[derive(serde::Serialize)]
struct QuantRowJson<'a> {
artifact: &'a str,
size_bytes: u64,
repo: &'a str,
#[serde(skip_serializing_if = "Option::is_none")]
bits: Option<f64>,
verification: VerificationJson,
#[serde(skip_serializing_if = "Option::is_none")]
verification_note: Option<&'a str>,
#[serde(skip_serializing_if = "Option::is_none")]
fit: Option<FitJson>,
}
#[derive(serde::Serialize)]
struct QuantsJson<'a> {
repo_id: &'a str,
artifacts: Vec<QuantRowJson<'a>>,
}
fn print_quants_json(
repo_id: &str,
rows: &[QuantArtifactRow],
fits: Option<&[FitVerdict]>,
) -> Result<(), FetchError> {
let artifacts: Vec<QuantRowJson<'_>> = rows
.iter()
.enumerate()
.map(|(i, row)| {
let (verification, verification_note) = match &row.verification {
discover::QuantVerification::Verified => (VerificationJson::Verified, None),
discover::QuantVerification::CheckFailed(reason) => {
(VerificationJson::CheckFailed, Some(reason.as_str()))
}
discover::QuantVerification::Unverified | _ => (VerificationJson::Unverified, None),
};
let fit = fits.and_then(|v| v.get(i)).map(|verdict| match verdict {
FitVerdict::FullGpu => FitJson {
fits: true,
resident_bytes: row.size,
n_cpu_moe: None,
moved_bytes: None,
reason: None,
},
FitVerdict::Offload {
n_cpu_moe,
moved_bytes,
resident_bytes,
} => FitJson {
fits: true,
resident_bytes: *resident_bytes,
n_cpu_moe: Some(*n_cpu_moe),
moved_bytes: Some(*moved_bytes),
reason: None,
},
FitVerdict::DoesNotFit { reason } => FitJson {
fits: false,
resident_bytes: row.size,
n_cpu_moe: None,
moved_bytes: None,
reason: Some(reason.clone()),
},
});
QuantRowJson {
artifact: row.artifact.as_str(),
size_bytes: row.size,
repo: row.repo.as_str(),
bits: row.bits,
verification,
verification_note,
fit,
}
})
.collect();
let serialized = serde_json::to_string_pretty(&QuantsJson { repo_id, artifacts })
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{serialized}");
Ok(())
}
fn run_quants(
repo_id: &str,
fits: Option<u64>,
reserve: Option<u64>,
token: Option<&str>,
json: bool,
) -> Result<(), FetchError> {
let owned_token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
eprintln!("Searching for quant siblings of {repo_id}...");
let client = hf_fetch_model::build_client(owned_token.as_deref())?;
let candidates = rt.block_on(discover::discover_quant_siblings(
repo_id,
owned_token.as_deref(),
&client,
))?;
let repo_count = candidates.len();
let verified_count = candidates
.iter()
.filter(|c| matches!(c.verification, discover::QuantVerification::Verified))
.count();
eprintln!("{repo_count} repos found, {verified_count} verified via GGUF backlink");
let rows = build_quant_rows(candidates);
if rows.is_empty() {
if json {
return print_quants_json(repo_id, &rows, None);
}
println!("No quant siblings found for {repo_id}.");
println!(
"Hint: quant repos are discovered by naming convention (<base>-<SCHEME>) — try `hf-fm search {repo_id}` for a broader view."
);
return Ok(());
}
let verdicts: Option<Vec<FitVerdict>> = if let Some(budget) = fits {
let reserve_bytes = reserve.unwrap_or(0);
let inspected = rows
.iter()
.filter(|r| trivial_fit_verdict(r, budget, reserve_bytes).is_none())
.count();
eprintln!("{inspected} inspected for offload plan");
Some(rt.block_on(compute_fit_verdicts_concurrent(
&rows,
budget,
reserve_bytes,
owned_token.as_deref(),
)))
} else {
None
};
if json {
return print_quants_json(repo_id, &rows, verdicts.as_deref());
}
print_quants_table(&rows, verdicts.as_deref());
Ok(())
}
fn print_model_card(card: &discover::ModelCardMetadata) {
println!();
if let Some(ref license) = card.license {
println!(" License: {license}");
}
if card.gated.is_gated() {
println!(
" Gated: {} (requires accepting terms on HF)",
card.gated
);
}
if let Some(ref pipeline) = card.pipeline_tag {
println!(" Pipeline: {pipeline}");
}
if let Some(ref library) = card.library_name {
println!(" Library: {library}");
}
if !card.tags.is_empty() {
println!(" Tags: {}", card.tags.join(", "));
}
if !card.languages.is_empty() {
println!(" Languages: {}", card.languages.join(", "));
}
}
fn run_info(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
json: bool,
max_lines: usize,
) -> Result<(), FetchError> {
if !repo_id.contains('/') {
return Err(FetchError::InvalidArgument(format!(
"invalid REPO_ID \"{repo_id}\": expected \"owner/model\" format \
(e.g., \"mistralai/Ministral-3-3B-Instruct-2512\")"
)));
}
let token_owned = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let card = rt.block_on(discover::fetch_model_card(repo_id))?;
let readme = rt.block_on(discover::fetch_readme(
repo_id,
revision,
token_owned.as_deref(),
))?;
if json {
return print_info_json(repo_id, &card, readme.as_deref());
}
println!(" Repo: {repo_id}");
print_model_card(&card);
if let Some(ref text) = readme {
let body = strip_yaml_front_matter(text);
println!();
println!(" README:");
if looks_like_default_template(body) {
println!(
" Note: README appears to be the HuggingFace default template (low information density)."
);
}
println!(" {}", "\u{2500}".repeat(70));
let lines: Vec<&str> = body.lines().collect();
let display_count = if max_lines == 0 {
lines.len()
} else {
lines.len().min(max_lines)
};
#[allow(clippy::indexing_slicing)]
for line in &lines[..display_count] {
println!(" {line}");
}
if display_count < lines.len() {
println!(
" ... ({} more lines, use --lines 0 for full output)",
lines.len().saturating_sub(display_count)
);
}
} else {
println!();
println!(" (no README.md found)");
}
Ok(())
}
#[derive(serde::Serialize)]
struct InfoResult {
repo_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
license: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pipeline_tag: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
library_name: Option<String>,
tags: Vec<String>,
languages: Vec<String>,
gated: String,
#[serde(skip_serializing_if = "Option::is_none")]
readme: Option<String>,
}
fn print_info_json(
repo_id: &str,
card: &discover::ModelCardMetadata,
readme: Option<&str>,
) -> Result<(), FetchError> {
let result = InfoResult {
repo_id: repo_id.to_owned(),
license: card.license.clone(),
pipeline_tag: card.pipeline_tag.clone(),
library_name: card.library_name.clone(),
tags: card.tags.clone(),
languages: card.languages.clone(),
gated: card.gated.to_string(),
readme: readme.map(str::to_owned),
};
let output = serde_json::to_string_pretty(&result)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
Ok(())
}
#[must_use]
fn strip_yaml_front_matter(text: &str) -> &str {
let trimmed = text.trim_start();
if !trimmed.starts_with("---") {
return text;
}
#[allow(clippy::indexing_slicing)]
let after_open = &trimmed[3..];
if let Some(close_pos) = after_open.find("\n---") {
let body_start = close_pos + 4; #[allow(clippy::indexing_slicing)]
let body = after_open[body_start..].trim_start_matches('\n');
return body.trim_start_matches('\r');
}
text
}
#[must_use]
fn looks_like_default_template(body: &str) -> bool {
let lines: Vec<&str> = body.lines().take(20).collect();
if lines.is_empty() {
return false;
}
if lines
.iter()
.any(|l| l.trim() == "# Model Card for Model ID")
{
return true;
}
let comment_count = lines
.iter()
.filter(|l| l.trim_start().starts_with("<!--"))
.count();
comment_count.saturating_mul(100) > lines.len().saturating_mul(30)
}
#[must_use]
fn format_header_line(header_size: u64, file_size: Option<u64>, is_headerless: bool) -> String {
if is_headerless {
match file_size {
Some(fs) => format!("Size: {}", format_size(fs)),
None => "Size: (size unknown)".to_owned(),
}
} else {
let hd = format_size(header_size);
match file_size {
Some(fs) => format!("Header: {hd} (JSON), {} total", format_size(fs)),
None => format!("Header: {hd} (JSON)"),
}
}
}
#[must_use]
fn format_metadata_lines(meta: &HashMap<String, String>) -> Vec<String> {
const TABULAR_THRESHOLD: usize = 6;
if meta.is_empty() {
return Vec::new();
}
let mut keys: Vec<&str> = meta.keys().map(String::as_str).collect();
keys.sort_unstable();
if keys.len() <= TABULAR_THRESHOLD {
let entries: Vec<String> = keys
.iter()
.map(|k| format!("{k}={}", meta.get(*k).map_or("", String::as_str)))
.collect();
return vec![format!("Metadata: {}", entries.join(", "))];
}
let mut lines: Vec<String> = Vec::with_capacity(keys.len() + 1);
lines.push("Metadata:".to_owned());
for k in &keys {
let v = meta.get(*k).map_or("", String::as_str);
if v.contains('\n') {
lines.push(format!(" {k}="));
for value_line in v.lines() {
lines.push(format!(" {value_line}"));
}
} else {
lines.push(format!(" {k}={v}"));
}
}
lines
}
#[must_use]
fn format_quant_lines(quant_info: Option<&inspect::QuantInfo>) -> Vec<String> {
let Some(q) = quant_info else {
return Vec::new();
};
vec![
format!("Format: {}", q.scheme),
format!(
"Size: {} stored -> {} (BF16)",
format_size(q.stored_bytes),
format_size(q.dequantized_bytes),
),
]
}
fn run_status_all(json: bool) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
let summaries = cache::cache_summary()?;
if json {
return print_status_all_json(&summaries, &cache_dir);
}
if summaries.is_empty() {
println!("No models found in local cache.");
return Ok(());
}
println!("Cache: {}\n", cache_dir.display());
let rw = summaries
.iter()
.map(|s| s.repo_id.len())
.max()
.unwrap_or(10)
.max(10); println!(
" {:<rw$} {:>5} {:>10} Status",
"Repository", "Files", "Size"
);
println!(" {:-<rw$} {:-<5} {:-<10} {:-<8}", "", "", "", "");
for s in &summaries {
let status_label = if s.has_partial { "PARTIAL" } else { "ok" };
println!(
" {:<rw$} {:>5} {:>10} {}",
s.repo_id,
s.file_count,
format_size(s.total_size),
status_label
);
}
println!(
"\n{} {} cached",
summaries.len(),
pluralize(summaries.len(), "model", "models")
);
Ok(())
}
fn resolve_du_arg(arg: &str) -> Result<String, FetchError> {
if arg.contains('/') {
return Ok(arg.to_owned());
}
if let Ok(n) = arg.parse::<usize>() {
let mut summaries = cache::cache_summary()?;
summaries.sort_by_key(|s| std::cmp::Reverse(s.total_size));
if n == 0 || n > summaries.len() {
return Err(FetchError::InvalidArgument(format!(
"index {n} is out of range (cache has {} {} — use 1..{})",
summaries.len(),
pluralize(summaries.len(), "repo", "repos"),
summaries.len()
)));
}
#[allow(clippy::indexing_slicing)]
return Ok(summaries[n - 1].repo_id.clone());
}
Err(FetchError::InvalidArgument(format!(
"\"{arg}\" is not a valid repo ID (expected \"org/model\") or numeric index"
)))
}
fn format_repo_size_cell(total_size: u64, gguf_size_range: Option<(u64, u64)>) -> String {
match gguf_size_range {
Some((min, max)) => format!("{} to {}", format_size(min), format_size(max)),
None => format_size(total_size),
}
}
fn run_du(age: bool, json: bool) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
let mut summaries = cache::cache_summary()?;
summaries.retain(|s| s.total_size > 0 || s.file_count > 0 || s.has_partial);
summaries.sort_by_key(|s| std::cmp::Reverse(s.total_size));
if json {
return print_du_json(&summaries, &cache_dir);
}
println!("Cache: {}\n", cache_dir.display());
if summaries.is_empty() {
println!("No models found in local cache.");
return Ok(());
}
let repo_width = summaries
.iter()
.map(|s| s.repo_id.len())
.max()
.unwrap_or(0)
.max(48);
let size_cells: Vec<String> = summaries
.iter()
.map(|s| format_repo_size_cell(s.total_size, s.gguf_size_range))
.collect();
let sw = size_cells.iter().map(String::len).max().unwrap_or(4).max(4);
if age {
println!(
" {:>3} {:>sw$} {:<repo_width$} {:>5} {:<15}",
"#", "SIZE", "REPO", "FILES", "AGE"
);
} else {
println!(
" {:>3} {:>sw$} {:<repo_width$} {:>5}",
"#", "SIZE", "REPO", "FILES"
);
}
let mut total_size: u64 = 0;
let mut total_files: usize = 0;
let mut any_partial = false;
let mut any_quant_alternatives = false;
for (i, (s, size_cell)) in summaries.iter().zip(size_cells.iter()).enumerate() {
total_size = total_size.saturating_add(s.total_size);
total_files = total_files.saturating_add(s.file_count);
any_quant_alternatives |= s.gguf_size_range.is_some();
let partial_marker = if s.has_partial {
any_partial = true;
" \u{25cf}"
} else {
""
};
if age {
let age_str = s
.last_modified
.map_or_else(|| "\u{2014}".to_owned(), format_age);
println!(
" {:>3} {:>sw$} {:<repo_width$} {:>5} {:<15}{}",
i + 1,
size_cell,
s.repo_id,
s.file_count,
age_str,
partial_marker,
);
} else {
println!(
" {:>3} {:>sw$} {:<repo_width$} {:>5}{}",
i + 1,
size_cell,
s.repo_id,
s.file_count,
partial_marker,
);
}
}
let rule_width = if age {
repo_width + sw + 36
} else {
repo_width + sw + 19
};
println!(" {}", "\u{2500}".repeat(rule_width));
println!(
" {:>sw$} total ({} {}, {} {})",
format_size(total_size),
summaries.len(),
pluralize(summaries.len(), "repo", "repos"),
total_files,
pluralize(total_files, "file", "files"),
);
if any_partial {
println!(" \u{25cf} = partial downloads");
}
if any_quant_alternatives {
println!(
" Note: a size range means that repo's cached `.gguf` files are \
mutually exclusive quant alternatives rather than shards of one \
file — you likely only need one of them (see `du <repo>` for the \
exact file sizes). The total above still reflects real bytes on \
disk across every cached file."
);
}
Ok(())
}
fn run_du_repo(repo_id: &str, json: bool) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
let files = cache::cache_repo_usage(repo_id)?;
let has_partial = cache::repo_has_partial(repo_id)?;
if json {
return print_du_repo_json(repo_id, &files, has_partial, &cache_dir);
}
println!("Cache: {}\n", cache_dir.display());
if files.is_empty() {
println!("No cached files found for {repo_id}.");
return Ok(());
}
println!(" {repo_id}:\n");
let fw = files
.iter()
.map(|f| f.filename.len())
.max()
.unwrap_or(4)
.max(4); let row_width = 3 + 2 + 10 + 2 + fw;
println!(" {:>3} {:>10} FILE", "#", "SIZE");
let mut total_size: u64 = 0;
for (i, f) in files.iter().enumerate() {
total_size = total_size.saturating_add(f.size);
println!(
" {:>3} {:>10} {}",
i + 1,
format_size(f.size),
f.filename
);
}
println!(" {}", "\u{2500}".repeat(row_width));
let sized: Vec<(&str, Option<u64>)> = files
.iter()
.map(|f| (f.filename.as_str(), Some(f.size)))
.collect();
if let Some((min, max)) = discover::gguf_size_range(sized) {
println!(
" {} to {} (mutually exclusive quants, {} {})",
format_size(min),
format_size(max),
files.len(),
pluralize(files.len(), "file", "files"),
);
} else {
println!(
" {:>10} total ({} {})",
format_size(total_size),
files.len(),
pluralize(files.len(), "file", "files"),
);
}
if has_partial {
println!("\n \u{25cf} partial downloads — run `hf-fm status {repo_id}` for details");
}
Ok(())
}
fn emit_json<T: serde::Serialize>(value: &T) -> Result<(), FetchError> {
let output = serde_json::to_string_pretty(value)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
Ok(())
}
fn system_time_to_unix(t: Option<std::time::SystemTime>) -> Option<u64> {
t.and_then(|st| st.duration_since(std::time::UNIX_EPOCH).ok())
.map(|d| d.as_secs())
}
#[derive(serde::Serialize)]
struct DuFileJson {
filename: String,
size: u64,
}
#[derive(serde::Serialize)]
struct DuRepoJson {
repo_id: String,
size: u64,
file_count: usize,
has_partial: bool,
last_modified: Option<u64>,
quant_alternatives: bool,
#[serde(skip_serializing_if = "Option::is_none")]
size_min: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
size_max: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
files: Option<Vec<DuFileJson>>,
}
#[derive(serde::Serialize)]
struct DuJson {
cache_dir: String,
repos: Vec<DuRepoJson>,
total_bytes: u64,
total_files: usize,
repo_count: usize,
}
#[derive(serde::Serialize)]
struct DuRepoDetailJson {
cache_dir: String,
repo_id: String,
files: Vec<DuFileJson>,
total_bytes: u64,
file_count: usize,
has_partial: bool,
quant_alternatives: bool,
#[serde(skip_serializing_if = "Option::is_none")]
size_min: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
size_max: Option<u64>,
}
fn print_du_json(
summaries: &[cache::CachedModelSummary],
cache_dir: &std::path::Path,
) -> Result<(), FetchError> {
let mut total_bytes: u64 = 0;
let mut total_files: usize = 0;
let mut repos: Vec<DuRepoJson> = Vec::with_capacity(summaries.len());
for s in summaries {
total_bytes = total_bytes.saturating_add(s.total_size);
total_files = total_files.saturating_add(s.file_count);
repos.push(DuRepoJson {
repo_id: s.repo_id.clone(),
size: s.total_size,
file_count: s.file_count,
has_partial: s.has_partial,
last_modified: system_time_to_unix(s.last_modified),
quant_alternatives: s.gguf_size_range.is_some(),
size_min: s.gguf_size_range.map(|(min, _)| min),
size_max: s.gguf_size_range.map(|(_, max)| max),
files: None,
});
}
let result = DuJson {
cache_dir: cache_dir.display().to_string(),
repo_count: repos.len(),
repos,
total_bytes,
total_files,
};
emit_json(&result)
}
fn print_du_tree_json(
repos: &[CacheTreeRepo],
cache_dir: &std::path::Path,
) -> Result<(), FetchError> {
let mut total_bytes: u64 = 0;
let mut total_files: usize = 0;
let mut out: Vec<DuRepoJson> = Vec::with_capacity(repos.len());
for r in repos {
total_bytes = total_bytes.saturating_add(r.total_size);
total_files = total_files.saturating_add(r.file_count);
let files = r
.files
.iter()
.map(|f| DuFileJson {
filename: f.filename.clone(),
size: f.size,
})
.collect();
out.push(DuRepoJson {
repo_id: r.repo_id.clone(),
size: r.total_size,
file_count: r.file_count,
has_partial: r.has_partial,
last_modified: system_time_to_unix(r.last_modified),
quant_alternatives: r.gguf_size_range.is_some(),
size_min: r.gguf_size_range.map(|(min, _)| min),
size_max: r.gguf_size_range.map(|(_, max)| max),
files: Some(files),
});
}
let result = DuJson {
cache_dir: cache_dir.display().to_string(),
repo_count: out.len(),
repos: out,
total_bytes,
total_files,
};
emit_json(&result)
}
fn print_du_repo_json(
repo_id: &str,
files: &[cache::CacheFileUsage],
has_partial: bool,
cache_dir: &std::path::Path,
) -> Result<(), FetchError> {
let mut total_bytes: u64 = 0;
let mut entries: Vec<DuFileJson> = Vec::with_capacity(files.len());
for f in files {
total_bytes = total_bytes.saturating_add(f.size);
entries.push(DuFileJson {
filename: f.filename.clone(),
size: f.size,
});
}
let sized: Vec<(&str, Option<u64>)> = files
.iter()
.map(|f| (f.filename.as_str(), Some(f.size)))
.collect();
let range = discover::gguf_size_range(sized);
let quant_alternatives = range.is_some();
let size_min = range.map(|(min, _)| min);
let size_max = range.map(|(_, max)| max);
let result = DuRepoDetailJson {
cache_dir: cache_dir.display().to_string(),
repo_id: repo_id.to_owned(),
file_count: entries.len(),
files: entries,
total_bytes,
has_partial,
quant_alternatives,
size_min,
size_max,
};
emit_json(&result)
}
const fn pluralize<'a>(n: usize, singular: &'a str, plural: &'a str) -> &'a str {
if n == 1 { singular } else { plural }
}
fn format_file_count(n: usize) -> String {
format!("({n} {})", pluralize(n, "file", "files"))
}
fn matches_filter(name: &str, pattern: &str) -> bool {
name.to_lowercase()
.contains(pattern.to_lowercase().as_str())
}
fn dot_filler(width: usize) -> String {
if width < 3 {
return " ".repeat(width);
}
let dots = width / 3;
let pad = width - 3 * dots;
let mut s = String::with_capacity(width);
for _ in 0..(pad + 2) {
s.push(' ');
}
for i in 0..dots {
if i > 0 {
s.push_str(" ");
}
s.push('.');
}
s
}
struct CacheTreeRepo {
repo_id: String,
total_size: u64,
file_count: usize,
has_partial: bool,
last_modified: Option<std::time::SystemTime>,
gguf_size_range: Option<(u64, u64)>,
files: Vec<CacheTreeFile>,
}
struct CacheTreeFile {
filename: String,
size: u64,
}
fn build_cache_tree() -> Result<Vec<CacheTreeRepo>, FetchError> {
let mut summaries = cache::cache_summary()?;
summaries.retain(|s| s.total_size > 0 || s.file_count > 0 || s.has_partial);
summaries.sort_by_key(|s| std::cmp::Reverse(s.total_size));
let mut repos: Vec<CacheTreeRepo> = Vec::with_capacity(summaries.len());
for s in summaries {
let usage = cache::cache_repo_usage(s.repo_id.as_str())?;
let files: Vec<CacheTreeFile> = usage
.into_iter()
.map(|f| CacheTreeFile {
filename: f.filename,
size: f.size,
})
.collect();
repos.push(CacheTreeRepo {
repo_id: s.repo_id,
total_size: s.total_size,
file_count: s.file_count,
has_partial: s.has_partial,
last_modified: s.last_modified,
gguf_size_range: s.gguf_size_range,
files,
});
}
Ok(repos)
}
const FILE_NAME_WIDTH_CAP: usize = 60;
struct CacheTreeWidths {
repo: usize,
size: usize,
files: usize,
age: usize,
file_name: usize,
file_size: usize,
}
impl CacheTreeWidths {
fn compute(repos: &[CacheTreeRepo], age: bool) -> Self {
let natural_repo = repos
.iter()
.map(|r| r.repo_id.len())
.max()
.unwrap_or(0)
.max(10);
let natural_size = repos
.iter()
.map(|r| format_repo_size_cell(r.total_size, r.gguf_size_range).len())
.max()
.unwrap_or(0)
.max(8);
let files = repos
.iter()
.map(|r| format_file_count(r.file_count).len())
.max()
.unwrap_or(0);
let age = if age {
repos
.iter()
.map(|r| r.last_modified.map_or(1, |t| format_age(t).len()))
.max()
.unwrap_or(0)
.max(15) } else {
0
};
let natural_file_name = repos
.iter()
.flat_map(|r| r.files.iter())
.map(|f| f.filename.len())
.max()
.unwrap_or(0)
.min(FILE_NAME_WIDTH_CAP);
let natural_file_size = repos
.iter()
.flat_map(|r| r.files.iter())
.map(|f| format_size(f.size).len())
.max()
.unwrap_or(0);
let size_start = (natural_repo + 8).max(natural_file_name + 12);
let repo = size_start - 8;
let file_name = size_start - 12;
let size = natural_size.max(natural_file_size);
Self {
repo,
size,
files,
age,
file_name,
file_size: size,
}
}
fn rule_width(&self) -> usize {
let mut w = 4 + self.repo + 2 + self.size + 2 + self.files;
if self.age > 0 {
w += 2 + self.age;
}
w
}
}
fn render_cache_tree(repos: &[CacheTreeRepo], widths: &CacheTreeWidths, age: bool) {
for (i, repo) in repos.iter().enumerate() {
let is_last = i + 1 == repos.len();
render_repo_node(repo, is_last, widths, age);
}
}
fn render_repo_node(repo: &CacheTreeRepo, is_last: bool, widths: &CacheTreeWidths, age: bool) {
let connector = if is_last { "└── " } else { "├── " };
let indent = if is_last { " " } else { "│ " };
let size_str = format_repo_size_cell(repo.total_size, repo.gguf_size_range);
let files_str = format_file_count(repo.file_count);
let partial_marker = if repo.has_partial { " \u{25cf}" } else { "" };
let filler = dot_filler((widths.repo + 2).saturating_sub(repo.repo_id.len()));
if age {
let age_str = repo
.last_modified
.map_or_else(|| "\u{2014}".to_owned(), format_age);
println!(
" {connector}{repo}{filler}{size:>sw$} {files:<fw$} {age_str:<aw$}{partial_marker}",
repo = repo.repo_id,
size = size_str,
files = files_str,
sw = widths.size,
fw = widths.files,
aw = widths.age,
);
} else {
println!(
" {connector}{repo}{filler}{size:>sw$} {files}{partial_marker}",
repo = repo.repo_id,
size = size_str,
files = files_str,
sw = widths.size,
);
}
render_file_leaves(&repo.files, indent, widths);
}
fn render_file_leaves(files: &[CacheTreeFile], indent: &str, widths: &CacheTreeWidths) {
if files.is_empty() {
return;
}
for (i, file) in files.iter().enumerate() {
let is_last = i + 1 == files.len();
let connector = if is_last { "└── " } else { "├── " };
println!(
" {indent}{connector}{name:<nw$} {size:>sw$}",
name = file.filename,
size = format_size(file.size),
nw = widths.file_name,
sw = widths.file_size,
);
}
}
fn run_du_tree(age: bool, json: bool) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
let repos = build_cache_tree()?;
if json {
return print_du_tree_json(&repos, &cache_dir);
}
println!("Cache: {}\n", cache_dir.display());
if repos.is_empty() {
println!("No models found in local cache.");
return Ok(());
}
let widths = CacheTreeWidths::compute(&repos, age);
render_cache_tree(&repos, &widths, age);
let total_size: u64 = repos
.iter()
.map(|r| r.total_size)
.fold(0_u64, u64::saturating_add);
let total_files: usize = repos
.iter()
.map(|r| r.file_count)
.fold(0_usize, usize::saturating_add);
let any_partial = repos.iter().any(|r| r.has_partial);
let any_quant_alternatives = repos.iter().any(|r| r.gguf_size_range.is_some());
println!("\n {}", "\u{2500}".repeat(widths.rule_width()));
println!(
" {:>10} total ({} {}, {} {})",
format_size(total_size),
repos.len(),
pluralize(repos.len(), "repo", "repos"),
total_files,
pluralize(total_files, "file", "files"),
);
if any_partial {
println!(" \u{25cf} = partial downloads");
}
if any_quant_alternatives {
println!(
" Note: a size range means that repo's cached `.gguf` files are \
mutually exclusive quant alternatives rather than shards of one \
file — you likely only need one of them. The total above still \
reflects real bytes on disk across every cached file."
);
}
Ok(())
}
fn run_cache_clean_partial(
repo_filter: Option<&str>,
yes: bool,
dry_run: bool,
) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
if !cache_dir.exists() {
println!("No HuggingFace cache found at {}", cache_dir.display());
return Ok(());
}
println!("Cache: {}\n", cache_dir.display());
let partials = cache::find_partial_files(repo_filter)?;
if partials.is_empty() {
println!("No partial downloads found.");
return Ok(());
}
let total_size: u64 = partials.iter().map(|p| p.size).sum();
if dry_run {
println!(
"Would remove {} {} ({}):",
partials.len(),
pluralize(partials.len(), "file", "files"),
format_size(total_size)
);
for p in &partials {
println!(" {}: {} ({})", p.repo_id, p.filename, format_size(p.size));
}
return Ok(());
}
println!(
"Found {} partial {}:",
partials.len(),
pluralize(partials.len(), "download", "downloads")
);
for p in &partials {
println!(" {}: {} ({})", p.repo_id, p.filename, format_size(p.size));
}
if !yes {
let prompt = format!(
"Clean {} {} ({})? [y/N]",
partials.len(),
pluralize(partials.len(), "file", "files"),
format_size(total_size)
);
if !confirm_prompt(prompt.as_str()) {
println!("Aborted.");
return Ok(());
}
}
for p in &partials {
std::fs::remove_file(&p.path).map_err(|e| FetchError::Io {
path: p.path.clone(),
source: e,
})?;
for sidecar in p.sidecar_paths() {
let _ = std::fs::remove_file(&sidecar);
}
}
println!(
"Removed {} {}. Freed {}.",
partials.len(),
pluralize(partials.len(), "file", "files"),
format_size(total_size)
);
Ok(())
}
fn delete_repo_dir(repo_id: &str) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_dir, repo_id);
if !repo_dir.exists() {
return Err(FetchError::InvalidArgument(format!(
"{repo_id} is not cached"
)));
}
let meta = std::fs::symlink_metadata(&repo_dir).map_err(|e| FetchError::Io {
path: repo_dir.clone(),
source: e,
})?;
if meta.file_type().is_symlink() {
return Err(FetchError::InvalidArgument(format!(
"{repo_id} cache entry is a symlink; refusing to delete"
)));
}
std::fs::remove_dir_all(&repo_dir).map_err(|e| FetchError::Io {
path: repo_dir.clone(),
source: e,
})?;
Ok(())
}
fn run_cache_delete(repo_id: &str, yes: bool) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
if !cache_dir.exists() {
println!("No HuggingFace cache found at {}", cache_dir.display());
return Ok(());
}
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_dir, repo_id);
if !repo_dir.exists() {
return Err(FetchError::InvalidArgument(format!(
"{repo_id} is not cached"
)));
}
let (file_count, size) = cache::repo_disk_usage(repo_id)?;
println!(
" {repo_id} ({}, {} {})",
format_size(size),
file_count,
pluralize(file_count, "file", "files")
);
if !yes && !confirm_prompt(" Delete? [y/N]") {
println!(" Aborted.");
return Ok(());
}
delete_repo_dir(repo_id)?;
println!(" Deleted. Freed {}.", format_size(size));
Ok(())
}
fn run_cache_path(repo_id: &str, revision: Option<&str>) -> Result<(), FetchError> {
let revision = revision.unwrap_or("main");
let cache_dir = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_dir, repo_id);
if !repo_dir.exists() {
return Err(FetchError::InvalidArgument(format!(
"{repo_id} is not cached"
)));
}
let commit_hash = cache::read_ref(&repo_dir, revision).ok_or_else(|| {
FetchError::InvalidArgument(format!(
"{repo_id} is cached but has no ref for \"{revision}\""
))
})?;
let snapshot_dir = hf_fetch_model::cache_layout::snapshot_dir(&repo_dir, commit_hash.as_str());
if !snapshot_dir.exists() {
return Err(FetchError::InvalidArgument(format!(
"snapshot directory for {repo_id} at revision \"{revision}\" does not exist"
)));
}
println!("{}", snapshot_dir.display());
Ok(())
}
fn format_verify_line(
filename: &str,
size: u64,
status: &cache::VerifyStatus,
fw: usize,
) -> String {
match status {
cache::VerifyStatus::Ok => format!(
" \u{2713} {filename:<fw$} {:>10} SHA256 OK",
format_size(size)
),
cache::VerifyStatus::Mismatch { expected, actual } => format!(
" \u{2717} {filename:<fw$} {:>10} SHA256 MISMATCH\n expected {expected}\n actual {actual}",
format_size(size)
),
cache::VerifyStatus::Skipped => format!(
" \u{2014} {filename:<fw$} {:>10} no LFS hash",
format_size(size)
),
cache::VerifyStatus::Missing => {
format!(" ! {filename:<fw$} {:>10} MISSING", format_size(size))
}
_ => format!(" ? {filename:<fw$} {:>10} UNKNOWN", format_size(size)),
}
}
#[allow(clippy::too_many_lines)]
fn run_cache_verify(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
) -> Result<(), FetchError> {
use indicatif::{ProgressBar, ProgressStyle};
use std::cell::Cell;
let cache_dir = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_dir, repo_id);
if !repo_dir.exists() {
return Err(FetchError::InvalidArgument(format!(
"{repo_id} is not cached"
)));
}
let resolved_token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let rev_display = revision.unwrap_or("main");
let commit_hash = cache::read_ref(&repo_dir, rev_display);
match commit_hash.as_deref() {
Some(hash) => println!("{repo_id} ({rev_display} @ {hash})"),
None => println!("{repo_id} ({rev_display}, ref not resolved)"),
}
println!("Cache: {}\n", cache_dir.display());
let bar = ProgressBar::new_spinner();
let style = ProgressStyle::with_template(" {spinner} {msg}")
.unwrap_or_else(|_| ProgressStyle::default_spinner())
.tick_chars("|/-\\ ");
bar.set_style(style);
bar.enable_steady_tick(Duration::from_millis(100));
let fw_cell: Cell<usize> = Cell::new(4);
let bar_ref = &bar;
let on_event = |event: cache::VerifyEvent<'_>| match event {
cache::VerifyEvent::Started {
total: _,
max_filename_len,
} => {
fw_cell.set(max_filename_len.max(4));
}
cache::VerifyEvent::FileStart {
index,
total,
filename,
size,
has_lfs,
} => {
let label = if has_lfs { "Verifying" } else { "Reading" };
bar_ref.set_message(format!(
"[{index}/{total}] {label} {filename} ({})",
format_size(size)
));
}
cache::VerifyEvent::FileComplete {
index: _,
total: _,
filename,
size,
status,
} => {
let line = format_verify_line(filename, size, status, fw_cell.get());
bar_ref.suspend(|| println!("{line}"));
}
_ => {}
};
let results = rt.block_on(cache::verify_cache_with_progress(
repo_id,
resolved_token.as_deref(),
revision,
on_event,
))?;
bar.finish_and_clear();
if results.is_empty() {
println!(" (no files found in remote repository)");
return Ok(());
}
let mut ok_count: usize = 0;
let mut mismatch_count: usize = 0;
let mut skipped_count: usize = 0;
let mut missing_count: usize = 0;
for r in &results {
match &r.status {
cache::VerifyStatus::Ok => ok_count += 1,
cache::VerifyStatus::Mismatch { .. } => mismatch_count += 1,
cache::VerifyStatus::Skipped => skipped_count += 1,
cache::VerifyStatus::Missing => missing_count += 1,
_ => {}
}
}
let total = results.len();
println!();
println!(
"{total} {}: {ok_count} SHA256 OK, {mismatch_count} mismatch, \
{skipped_count} skipped, {missing_count} missing",
pluralize(total, "file", "files")
);
if mismatch_count > 0 {
return Err(FetchError::InvalidArgument(format!(
"{mismatch_count} {} failed SHA256 verification",
pluralize(mismatch_count, "file", "files")
)));
}
Ok(())
}
const FRESH_PARTIAL_WINDOW_SECS: u64 = 60 * 60;
#[derive(Debug, Clone)]
struct EvictionEntry {
repo_id: String,
size: u64,
last_modified: Option<std::time::SystemTime>,
}
impl From<&cache::CachedModelSummary> for EvictionEntry {
fn from(s: &cache::CachedModelSummary) -> Self {
Self {
repo_id: s.repo_id.clone(),
size: s.total_size,
last_modified: s.last_modified,
}
}
}
struct GcCriteria {
older_than_secs: Option<u64>,
max_size: Option<u64>,
except: HashSet<String>,
}
struct GcPlan {
evict: Vec<EvictionEntry>,
protected: Vec<EvictionEntry>,
kept: Vec<EvictionEntry>,
skipped_partials: Vec<EvictionEntry>,
size_before: u64,
size_after: u64,
budget_shortfall: bool,
}
#[must_use]
fn select_age_evictions(
summaries: &[cache::CachedModelSummary],
threshold_secs: u64,
now: std::time::SystemTime,
) -> HashSet<&str> {
summaries
.iter()
.filter_map(|s| {
let mtime = s.last_modified?;
let elapsed = now.duration_since(mtime).ok()?;
(elapsed.as_secs() >= threshold_secs).then_some(s.repo_id.as_str())
})
.collect()
}
#[must_use]
fn compute_gc_plan(
summaries: &[cache::CachedModelSummary],
criteria: &GcCriteria,
now: std::time::SystemTime,
list_kept: bool,
) -> GcPlan {
let size_before: u64 = summaries
.iter()
.map(|s| s.total_size)
.fold(0_u64, u64::saturating_add);
let mut protected: Vec<EvictionEntry> = Vec::new();
let mut skipped_partials: Vec<EvictionEntry> = Vec::new();
let mut eligible: Vec<&cache::CachedModelSummary> = Vec::new();
for s in summaries {
if criteria.except.contains(s.repo_id.as_str()) {
protected.push(EvictionEntry::from(s));
continue;
}
let is_fresh_partial = s.has_partial
&& s.last_modified.is_some_and(|m| {
now.duration_since(m)
.is_ok_and(|d| d.as_secs() < FRESH_PARTIAL_WINDOW_SECS)
});
if is_fresh_partial {
skipped_partials.push(EvictionEntry::from(s));
} else {
eligible.push(s);
}
}
let age_set: HashSet<&str> = match criteria.older_than_secs {
Some(threshold) => select_age_evictions(summaries, threshold, now),
None => HashSet::new(),
};
let (mut to_evict, mut still_eligible): (
Vec<&cache::CachedModelSummary>,
Vec<&cache::CachedModelSummary>,
) = eligible
.into_iter()
.partition(|s| age_set.contains(s.repo_id.as_str()));
let mut budget_shortfall = false;
if let Some(max_size) = criteria.max_size {
let already_evicting: u64 = to_evict
.iter()
.map(|s| s.total_size)
.fold(0_u64, u64::saturating_add);
let mut remaining_after = size_before.saturating_sub(already_evicting);
if remaining_after > max_size {
still_eligible.sort_by(|a, b| {
(a.last_modified, a.repo_id.as_str()).cmp(&(b.last_modified, b.repo_id.as_str()))
});
let mut split_idx = 0_usize;
for s in &still_eligible {
if remaining_after <= max_size {
break;
}
to_evict.push(s);
remaining_after = remaining_after.saturating_sub(s.total_size);
split_idx = split_idx.saturating_add(1);
}
budget_shortfall = remaining_after > max_size;
still_eligible.drain(0..split_idx);
}
}
to_evict.sort_by(|a, b| {
(a.last_modified, a.repo_id.as_str()).cmp(&(b.last_modified, b.repo_id.as_str()))
});
let evict: Vec<EvictionEntry> = to_evict.iter().map(|s| EvictionEntry::from(*s)).collect();
let evicted_size: u64 = evict
.iter()
.map(|e| e.size)
.fold(0_u64, u64::saturating_add);
let size_after = size_before.saturating_sub(evicted_size);
let kept = if list_kept {
still_eligible
.iter()
.map(|s| EvictionEntry::from(*s))
.collect()
} else {
Vec::new()
};
GcPlan {
evict,
protected,
kept,
skipped_partials,
size_before,
size_after,
budget_shortfall,
}
}
fn render_eviction_table(entries: &[EvictionEntry]) {
let name_width = entries.iter().map(|e| e.repo_id.len()).max().unwrap_or(0);
let size_width = entries
.iter()
.map(|e| format_size(e.size).len())
.max()
.unwrap_or(0);
for entry in entries {
let age_str = entry
.last_modified
.map_or_else(|| "\u{2014}".to_owned(), format_age);
println!(
" {:<nw$} {:>sw$} {age_str}",
entry.repo_id,
format_size(entry.size),
nw = name_width,
sw = size_width,
);
}
}
fn render_gc_plan_preview(plan: &GcPlan) {
if !plan.evict.is_empty() {
println!("Will remove:");
render_eviction_table(&plan.evict);
println!();
}
if !plan.skipped_partials.is_empty() {
println!("Skipped (active partial downloads):");
for entry in &plan.skipped_partials {
println!(" {}", entry.repo_id);
}
println!();
}
if !plan.protected.is_empty() {
println!("Protected by --except:");
for entry in &plan.protected {
println!(" {}", entry.repo_id);
}
println!();
}
if !plan.kept.is_empty() {
println!("Keep:");
render_eviction_table(&plan.kept);
println!();
}
let freed = plan.size_before.saturating_sub(plan.size_after);
println!(
"Cache: {} \u{2192} {} (free {})",
format_size(plan.size_before),
format_size(plan.size_after),
format_size(freed)
);
if plan.budget_shortfall {
eprintln!("warning: --max-size budget cannot be reached; protected repos exceed the cap");
}
}
fn run_cache_gc(
older_than_days: Option<u64>,
max_size: Option<u64>,
except: Vec<String>,
dry_run: bool,
yes: bool,
list_kept: bool,
) -> Result<(), FetchError> {
let cache_dir = cache::hf_cache_dir()?;
if !cache_dir.exists() {
println!("No HuggingFace cache found at {}", cache_dir.display());
return Ok(());
}
println!("Cache: {}\n", cache_dir.display());
let summaries = cache::cache_summary()?;
if summaries.is_empty() {
println!("No models in cache.");
return Ok(());
}
let known_ids: HashSet<&str> = summaries.iter().map(|s| s.repo_id.as_str()).collect();
for repo in &except {
if !known_ids.contains(repo.as_str()) {
eprintln!("warning: --except {repo:?} is not a cached repo; ignoring");
}
}
let criteria = GcCriteria {
older_than_secs: older_than_days.map(|d| d.saturating_mul(86_400)),
max_size,
except: except.into_iter().collect(),
};
let plan = compute_gc_plan(
&summaries,
&criteria,
std::time::SystemTime::now(),
list_kept,
);
if !plan.skipped_partials.is_empty() {
let n = plan.skipped_partials.len();
eprintln!(
"note: {n} {} skipped (active partial {}); run `hf-fm cache clean-partial` first",
pluralize(n, "repo", "repos"),
pluralize(n, "download", "downloads")
);
}
if plan.evict.is_empty() {
println!("No repos matched eviction criteria.");
return Ok(());
}
render_gc_plan_preview(&plan);
if dry_run {
return Ok(());
}
if !yes && !confirm_prompt("\nProceed? [y/N]") {
println!("Aborted.");
return Ok(());
}
let mut failures: Vec<(String, FetchError)> = Vec::new();
let mut freed: u64 = 0;
for entry in &plan.evict {
match delete_repo_dir(entry.repo_id.as_str()) {
Ok(()) => freed = freed.saturating_add(entry.size),
Err(e) => {
eprintln!("error: failed to delete {}: {e}", entry.repo_id);
failures.push((entry.repo_id.clone(), e));
}
}
}
let success_count = plan.evict.len().saturating_sub(failures.len());
println!(
"Removed {success_count} {}. Freed {}.",
pluralize(success_count, "repo", "repos"),
format_size(freed)
);
if !failures.is_empty() {
let n = failures.len();
return Err(FetchError::InvalidArgument(format!(
"{n} {} failed to delete during gc",
pluralize(n, "repo", "repos")
)));
}
Ok(())
}
fn confirm_prompt(message: &str) -> bool {
eprint!("{message} ");
let mut input = String::new();
std::io::stdin().read_line(&mut input).is_ok() && input.trim().eq_ignore_ascii_case("y")
}
fn collect_repo_tensors(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Result<HashMap<String, inspect::TensorInfo>, FetchError> {
let results: Vec<(String, inspect::SafetensorsHeaderInfo)> = if cached {
inspect::inspect_repo_safetensors_cached(repo_id, revision)?
} else {
let token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let remote_results = rt.block_on(inspect::inspect_repo_safetensors(
repo_id,
token.as_deref(),
revision,
))?;
remote_results
.into_iter()
.map(|(name, info, _source)| (name, info))
.collect()
};
let mut tensors = HashMap::new();
for (_filename, info) in results {
for t in info.tensors {
tensors.insert(t.name.clone(), t);
}
}
Ok(tensors)
}
#[allow(
clippy::too_many_arguments,
clippy::fn_params_excessive_bools,
clippy::too_many_lines
)]
fn run_diff(
repo_a: &str,
repo_b: &str,
revision_a: Option<&str>,
revision_b: Option<&str>,
token: Option<&str>,
cached: bool,
filter: Option<&str>,
summary: bool,
dtypes: bool,
collapse: bool,
limit: Option<usize>,
json: bool,
) -> Result<(), FetchError> {
let tensors_a = collect_repo_tensors(repo_a, revision_a, token, cached)
.map_err(|e| enrich_gated_content_error(e, repo_a, token))?;
let tensors_b = collect_repo_tensors(repo_b, revision_b, token, cached)
.map_err(|e| enrich_gated_content_error(e, repo_b, token))?;
if !json {
let empty: Vec<&str> = [(repo_a, &tensors_a), (repo_b, &tensors_b)]
.into_iter()
.filter_map(|(repo, tensors)| tensors.is_empty().then_some(repo))
.collect();
if !empty.is_empty() {
for repo in empty {
println!("No .safetensors files found in {repo}.");
println!("Hint: use `hf-fm list-files {repo}` to see available file types");
}
return Ok(());
}
}
let mut all_names: Vec<&str> = tensors_a
.keys()
.chain(tensors_b.keys())
.map(String::as_str)
.collect::<BTreeSet<&str>>()
.into_iter()
.collect();
if let Some(pattern) = filter {
all_names.retain(|name| matches_filter(name, pattern));
}
let mut only_a: Vec<&str> = Vec::new();
let mut only_b: Vec<&str> = Vec::new();
let mut differ: Vec<&str> = Vec::new();
let mut matching: Vec<&str> = Vec::new();
for name in &all_names {
match (tensors_a.get(*name), tensors_b.get(*name)) {
(Some(_), None) => only_a.push(name),
(None, Some(_)) => only_b.push(name),
(Some(a), Some(b)) => {
if a.dtype == b.dtype && a.shape == b.shape {
matching.push(name);
} else {
differ.push(name);
}
}
(None, None) => {} }
}
let total_a = if filter.is_some() {
only_a.len() + differ.len() + matching.len()
} else {
tensors_a.len()
};
let total_b = if filter.is_some() {
only_b.len() + differ.len() + matching.len()
} else {
tensors_b.len()
};
if json {
return print_diff_json(
repo_a, repo_b, &tensors_a, &tensors_b, &only_a, &only_b, &differ, &matching, filter,
dtypes, collapse, limit,
);
}
if dtypes {
print_diff_dtypes(repo_a, repo_b, &tensors_a, &tensors_b, filter);
return Ok(());
}
if collapse {
print_diff_collapse(
repo_a, repo_b, &tensors_a, &tensors_b, &only_a, &only_b, &differ, filter, limit,
);
return Ok(());
}
print_diff_repo_header(repo_a, repo_b);
if !summary {
let nw = only_a
.iter()
.chain(only_b.iter())
.map(|n| n.len())
.max()
.unwrap_or(0);
let cap = limit.unwrap_or(usize::MAX);
println!();
if !only_a.is_empty() {
let label = if only_a.len() == 1 {
"tensor"
} else {
"tensors"
};
println!(" Only in A ({} {label}):", only_a.len());
for name in only_a.iter().take(cap) {
if let Some(t) = tensors_a.get(*name) {
let shape_str = format!("{:?}", t.shape);
println!(" {name:<nw$} {:<8} {shape_str}", t.dtype);
}
}
print_truncation_note(cap, only_a.len());
println!();
}
if !only_b.is_empty() {
let label = if only_b.len() == 1 {
"tensor"
} else {
"tensors"
};
println!(" Only in B ({} {label}):", only_b.len());
for name in only_b.iter().take(cap) {
if let Some(t) = tensors_b.get(*name) {
let shape_str = format!("{:?}", t.shape);
println!(" {name:<nw$} {:<8} {shape_str}", t.dtype);
}
}
print_truncation_note(cap, only_b.len());
println!();
}
if !differ.is_empty() {
let label = if differ.len() == 1 {
"tensor"
} else {
"tensors"
};
println!(" Dtype/shape differences ({} {label}):", differ.len());
for name in differ.iter().take(cap) {
if let Some((a, b)) = tensors_a.get(*name).zip(tensors_b.get(*name)) {
let shape_a = format!("{:?}", a.shape);
let shape_b = format!("{:?}", b.shape);
println!(" {name}");
println!(" A: {:<8} {shape_a}", a.dtype);
println!(" B: {:<8} {shape_b}", b.dtype);
}
}
print_truncation_note(cap, differ.len());
println!();
}
let match_label = if matching.len() == 1 {
"tensor"
} else {
"tensors"
};
println!(" Matching: {} {match_label} identical", matching.len());
}
println!(" {}", "\u{2500}".repeat(70));
print!(
" A: {} {} | B: {} {} | only-A: {} | only-B: {} | differ: {} | match: {}",
total_a,
pluralize(total_a, "tensor", "tensors"),
total_b,
pluralize(total_b, "tensor", "tensors"),
only_a.len(),
only_b.len(),
differ.len(),
matching.len(),
);
if let Some(pattern) = filter {
println!(" (filter: {pattern:?})");
} else {
println!();
}
Ok(())
}
#[derive(serde::Serialize)]
struct DiffTensorEntry {
name: String,
#[serde(skip_serializing_if = "Option::is_none")]
a: Option<DiffTensorSide>,
#[serde(skip_serializing_if = "Option::is_none")]
b: Option<DiffTensorSide>,
}
#[derive(serde::Serialize)]
struct DiffTensorSide {
dtype: String,
shape: Vec<usize>,
byte_count: u64,
}
#[derive(serde::Serialize)]
struct DiffDtypeGroup {
dtype: String,
tensors: usize,
params: u64,
bytes: u64,
}
#[derive(serde::Serialize)]
struct DtypeHistograms {
a: Vec<DiffDtypeGroup>,
b: Vec<DiffDtypeGroup>,
}
#[derive(serde::Serialize)]
struct DiffCollapsed {
only_a: Vec<DiffCollapseGroup>,
only_b: Vec<DiffCollapseGroup>,
differ: Vec<DiffCollapseGroup>,
}
#[derive(serde::Serialize)]
struct DiffTruncation {
only_a: TruncationInfo,
only_b: TruncationInfo,
differ: TruncationInfo,
}
#[derive(serde::Serialize)]
struct DiffResult {
repo_a: String,
repo_b: String,
only_a: Vec<DiffTensorEntry>,
only_b: Vec<DiffTensorEntry>,
differ: Vec<DiffTensorEntry>,
matching_count: usize,
#[serde(skip_serializing_if = "Option::is_none")]
filter: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
dtype_histograms: Option<DtypeHistograms>,
#[serde(skip_serializing_if = "Option::is_none")]
collapsed: Option<DiffCollapsed>,
#[serde(skip_serializing_if = "Option::is_none")]
truncated: Option<DiffTruncation>,
}
#[allow(clippy::too_many_arguments)]
fn print_diff_json(
repo_a: &str,
repo_b: &str,
tensors_a: &HashMap<String, inspect::TensorInfo>,
tensors_b: &HashMap<String, inspect::TensorInfo>,
only_a: &[&str],
only_b: &[&str],
differ: &[&str],
matching: &[&str],
filter: Option<&str>,
dtypes: bool,
collapse: bool,
limit: Option<usize>,
) -> Result<(), FetchError> {
let make_entry = |name: &str,
a: Option<&inspect::TensorInfo>,
b: Option<&inspect::TensorInfo>|
-> DiffTensorEntry {
DiffTensorEntry {
name: name.to_owned(),
a: a.map(|t| DiffTensorSide {
dtype: t.dtype.clone(),
shape: t.shape.clone(),
byte_count: t.byte_len(),
}),
b: b.map(|t| DiffTensorSide {
dtype: t.dtype.clone(),
shape: t.shape.clone(),
byte_count: t.byte_len(),
}),
}
};
let dtype_histograms = if dtypes {
let (rows_a, rows_b) = aggregate_diff_dtypes(tensors_a, tensors_b, filter);
Some(DtypeHistograms {
a: rows_a,
b: rows_b,
})
} else {
None
};
let collapsed = if collapse {
Some(DiffCollapsed {
only_a: aggregate_diff_collapse(only_a, tensors_a, None),
only_b: aggregate_diff_collapse(only_b, tensors_b, None),
differ: aggregate_diff_collapse(differ, tensors_a, Some(tensors_b)),
})
} else {
None
};
let cap = limit.unwrap_or(usize::MAX);
let truncated = limit
.is_some_and(|n| only_a.len() > n || only_b.len() > n || differ.len() > n)
.then(|| DiffTruncation {
only_a: TruncationInfo {
shown: only_a.len().min(cap),
total: only_a.len(),
},
only_b: TruncationInfo {
shown: only_b.len().min(cap),
total: only_b.len(),
},
differ: TruncationInfo {
shown: differ.len().min(cap),
total: differ.len(),
},
});
let result = DiffResult {
repo_a: repo_a.to_owned(),
repo_b: repo_b.to_owned(),
only_a: only_a
.iter()
.take(cap)
.map(|n| make_entry(n, tensors_a.get(*n), None))
.collect(),
only_b: only_b
.iter()
.take(cap)
.map(|n| make_entry(n, None, tensors_b.get(*n)))
.collect(),
differ: differ
.iter()
.take(cap)
.map(|n| make_entry(n, tensors_a.get(*n), tensors_b.get(*n)))
.collect(),
matching_count: matching.len(),
filter: filter.map(str::to_owned),
dtype_histograms,
collapsed,
truncated,
};
emit_json(&result)
}
#[allow(dead_code)] const DIFF_DTYPES_DESIGN_WIDTH: usize = 75;
struct DiffDtypesColumnWidths {
dtype: usize,
a_tensors: usize,
a_size: usize,
b_tensors: usize,
b_size: usize,
delta: usize,
}
impl DiffDtypesColumnWidths {
fn total_width(&self) -> usize {
2 + self.dtype
+ 2
+ self.a_tensors
+ 2
+ self.a_size
+ 2
+ self.b_tensors
+ 2
+ self.b_size
+ 2
+ self.delta
}
}
fn format_size_delta(delta: i64) -> String {
match delta.cmp(&0) {
std::cmp::Ordering::Less => format!("-{}", format_size(delta.unsigned_abs())),
std::cmp::Ordering::Greater => format!("+{}", format_size(delta.unsigned_abs())),
std::cmp::Ordering::Equal => format_size(0),
}
}
fn format_count_delta(delta: i64) -> String {
if delta == 0 {
"0".to_owned()
} else {
format!("{delta:+}")
}
}
fn column_width<'a>(header: &str, cells: impl IntoIterator<Item = &'a str>) -> usize {
cells.into_iter().fold(header.chars().count(), |w, cell| {
w.max(cell.chars().count())
})
}
fn print_truncation_note(cap: usize, total: usize) {
if total > cap {
println!(" \u{2026} showing {cap} of {total} (limit {cap})");
}
}
fn print_diff_repo_header(repo_a: &str, repo_b: &str) {
println!(" A: {repo_a}");
println!(" B: {repo_b}");
}
fn aggregate_diff_dtypes(
tensors_a: &HashMap<String, inspect::TensorInfo>,
tensors_b: &HashMap<String, inspect::TensorInfo>,
filter: Option<&str>,
) -> (Vec<DiffDtypeGroup>, Vec<DiffDtypeGroup>) {
let collect = |map: &HashMap<String, inspect::TensorInfo>| -> Vec<DiffDtypeGroup> {
let filtered = map
.iter()
.filter(|(name, _)| filter.is_none_or(|p| matches_filter(name, p)))
.map(|(_, t)| t);
let mut rows: Vec<DiffDtypeGroup> = compute_dtype_groups(filtered)
.into_iter()
.map(|(dtype, tensors, params, bytes)| DiffDtypeGroup {
dtype: dtype.to_owned(),
tensors,
params,
bytes,
})
.collect();
rows.sort_by_key(|r| std::cmp::Reverse(r.bytes));
rows
};
(collect(tensors_a), collect(tensors_b))
}
fn diff_dtypes_column_widths(
rows_a: &[DiffDtypeGroup],
rows_b: &[DiffDtypeGroup],
) -> DiffDtypesColumnWidths {
let mut dtype_w = "Dtype".chars().count();
let mut a_tensors_w = "A Tensors".chars().count();
let mut a_size_w = "A Size".chars().count();
let mut b_tensors_w = "B Tensors".chars().count();
let mut b_size_w = "B Size".chars().count();
let mut delta_w = "\u{0394} Size".chars().count();
let mut all_dtypes: BTreeSet<&str> = BTreeSet::new();
for r in rows_a {
all_dtypes.insert(r.dtype.as_str());
}
for r in rows_b {
all_dtypes.insert(r.dtype.as_str());
}
for dtype in &all_dtypes {
dtype_w = dtype_w.max(dtype.chars().count());
let a = rows_a.iter().find(|r| r.dtype == *dtype);
let b = rows_b.iter().find(|r| r.dtype == *dtype);
let a_tensors_str = a.map_or_else(|| "\u{2014}".to_owned(), |r| r.tensors.to_string());
let b_tensors_str = b.map_or_else(|| "\u{2014}".to_owned(), |r| r.tensors.to_string());
let a_size_str = a.map_or_else(|| "\u{2014}".to_owned(), |r| format_size(r.bytes));
let b_size_str = b.map_or_else(|| "\u{2014}".to_owned(), |r| format_size(r.bytes));
a_tensors_w = a_tensors_w.max(a_tensors_str.chars().count());
b_tensors_w = b_tensors_w.max(b_tensors_str.chars().count());
a_size_w = a_size_w.max(a_size_str.chars().count());
b_size_w = b_size_w.max(b_size_str.chars().count());
let a_bytes_i = a.map_or(0_i64, |r| i64::try_from(r.bytes).unwrap_or(i64::MAX));
let b_bytes_i = b.map_or(0_i64, |r| i64::try_from(r.bytes).unwrap_or(i64::MAX));
let delta = b_bytes_i.saturating_sub(a_bytes_i);
delta_w = delta_w.max(format_size_delta(delta).chars().count());
}
DiffDtypesColumnWidths {
dtype: dtype_w,
a_tensors: a_tensors_w,
a_size: a_size_w,
b_tensors: b_tensors_w,
b_size: b_size_w,
delta: delta_w,
}
}
#[allow(clippy::similar_names)]
fn print_diff_dtypes(
repo_a: &str,
repo_b: &str,
tensors_a: &HashMap<String, inspect::TensorInfo>,
tensors_b: &HashMap<String, inspect::TensorInfo>,
filter: Option<&str>,
) {
let (rows_a, rows_b) = aggregate_diff_dtypes(tensors_a, tensors_b, filter);
let w = diff_dtypes_column_widths(&rows_a, &rows_b);
let dw = w.dtype;
let atw = w.a_tensors;
let asw = w.a_size;
let btw = w.b_tensors;
let bsw = w.b_size;
let dlw = w.delta;
print_diff_repo_header(repo_a, repo_b);
println!();
println!(
" {:<dw$} {:>atw$} {:>asw$} {:>btw$} {:>bsw$} {:>dlw$}",
"Dtype", "A Tensors", "A Size", "B Tensors", "B Size", "\u{0394} Size",
);
let mut all_dtypes: Vec<&str> = rows_a
.iter()
.chain(rows_b.iter())
.map(|r| r.dtype.as_str())
.collect::<BTreeSet<&str>>()
.into_iter()
.collect();
all_dtypes.sort_by_key(|dtype| {
let a_bytes = rows_a
.iter()
.find(|r| r.dtype == *dtype)
.map_or(0, |r| r.bytes);
let b_bytes = rows_b
.iter()
.find(|r| r.dtype == *dtype)
.map_or(0, |r| r.bytes);
std::cmp::Reverse(a_bytes.max(b_bytes))
});
for dtype in &all_dtypes {
let a = rows_a.iter().find(|r| r.dtype == *dtype);
let b = rows_b.iter().find(|r| r.dtype == *dtype);
let a_tensors = a.map_or_else(|| "\u{2014}".to_owned(), |r| r.tensors.to_string());
let b_tensors = b.map_or_else(|| "\u{2014}".to_owned(), |r| r.tensors.to_string());
let a_size = a.map_or_else(|| "\u{2014}".to_owned(), |r| format_size(r.bytes));
let b_size = b.map_or_else(|| "\u{2014}".to_owned(), |r| format_size(r.bytes));
let a_bytes_i = a.map_or(0_i64, |r| i64::try_from(r.bytes).unwrap_or(i64::MAX));
let b_bytes_i = b.map_or(0_i64, |r| i64::try_from(r.bytes).unwrap_or(i64::MAX));
let delta_str = format_size_delta(b_bytes_i.saturating_sub(a_bytes_i));
println!(
" {dtype:<dw$} {a_tensors:>atw$} {a_size:>asw$} {b_tensors:>btw$} {b_size:>bsw$} {delta_str:>dlw$}",
);
}
println!(" {}", "\u{2500}".repeat(w.total_width().saturating_sub(2)));
let total_a_tensors: usize = rows_a.iter().map(|r| r.tensors).sum();
let total_b_tensors: usize = rows_b.iter().map(|r| r.tensors).sum();
let total_a_bytes: u64 = rows_a.iter().map(|r| r.bytes).sum();
let total_b_bytes: u64 = rows_b.iter().map(|r| r.bytes).sum();
let delta_tensors = i64::try_from(total_b_tensors)
.unwrap_or(i64::MAX)
.saturating_sub(i64::try_from(total_a_tensors).unwrap_or(i64::MAX));
let delta_bytes = i64::try_from(total_b_bytes)
.unwrap_or(i64::MAX)
.saturating_sub(i64::try_from(total_a_bytes).unwrap_or(i64::MAX));
let footer_core = format!(
" A: {} {}, {} | B: {} {}, {} | \u{0394}: {} tensors, {}",
total_a_tensors,
pluralize(total_a_tensors, "tensor", "tensors"),
format_size(total_a_bytes),
total_b_tensors,
pluralize(total_b_tensors, "tensor", "tensors"),
format_size(total_b_bytes),
format_count_delta(delta_tensors),
format_size_delta(delta_bytes),
);
if let Some(p) = filter {
println!("{footer_core} (filter: {p:?})");
} else {
println!("{footer_core}");
}
}
fn collapse_numeric_segments(name: &str) -> String {
let mut out = String::with_capacity(name.len());
let mut chars = name.chars().peekable();
while let Some(c) = chars.next() {
if c.is_ascii_digit() {
out.push_str("{N}");
while chars.next_if(char::is_ascii_digit).is_some() {}
} else {
out.push(c);
}
}
out
}
#[derive(serde::Serialize)]
struct DiffCollapseGroup {
pattern: String,
tensors: usize,
bytes: u64,
#[serde(skip_serializing_if = "Option::is_none")]
bytes_b: Option<u64>,
}
fn sum_named_tensor_bytes(members: &[&str], tensors: &HashMap<String, inspect::TensorInfo>) -> u64 {
members
.iter()
.filter_map(|n| tensors.get(*n))
.map(inspect::TensorInfo::byte_len)
.sum()
}
fn aggregate_diff_collapse(
names: &[&str],
primary: &HashMap<String, inspect::TensorInfo>,
secondary: Option<&HashMap<String, inspect::TensorInfo>>,
) -> Vec<DiffCollapseGroup> {
let mut by_pattern: HashMap<String, Vec<&str>> = HashMap::new();
for name in names {
by_pattern
.entry(collapse_numeric_segments(name))
.or_default()
.push(name);
}
let mut rows: Vec<DiffCollapseGroup> = by_pattern
.into_iter()
.map(|(pattern, members)| {
let tensors = members.len();
let bytes = sum_named_tensor_bytes(&members, primary);
let bytes_b = secondary.map(|s| sum_named_tensor_bytes(&members, s));
DiffCollapseGroup {
pattern,
tensors,
bytes,
bytes_b,
}
})
.collect();
rows.sort_by(|a, b| {
b.bytes
.cmp(&a.bytes)
.then_with(|| a.pattern.cmp(&b.pattern))
});
rows
}
fn print_diff_collapse_section(
label: &str,
rows: &[DiffCollapseGroup],
cap: usize,
total_tensors: usize,
) {
println!(
" {label} ({total_tensors} {}, {} {}):",
pluralize(total_tensors, "tensor", "tensors"),
rows.len(),
pluralize(rows.len(), "pattern", "patterns"),
);
let cells: Vec<(String, String)> = rows
.iter()
.map(|r| (r.tensors.to_string(), format_size(r.bytes)))
.collect();
let pw = column_width("Pattern", rows.iter().map(|r| r.pattern.as_str()));
let tw = column_width("Tensors", cells.iter().map(|(t, _)| t.as_str()));
let bw = column_width("Bytes", cells.iter().map(|(_, b)| b.as_str()));
println!(" {:<pw$} {:>tw$} {:>bw$}", "Pattern", "Tensors", "Bytes");
for (r, (tensors, bytes)) in rows.iter().zip(&cells).take(cap) {
println!(" {:<pw$} {tensors:>tw$} {bytes:>bw$}", r.pattern);
}
print_truncation_note(cap, rows.len());
println!();
}
fn print_diff_collapse_differ(rows: &[DiffCollapseGroup], cap: usize, total_tensors: usize) {
println!(
" Dtype/shape differences ({total_tensors} {}, {} {}):",
pluralize(total_tensors, "tensor", "tensors"),
rows.len(),
pluralize(rows.len(), "pattern", "patterns"),
);
let cells: Vec<(String, String, String, String)> = rows
.iter()
.map(|r| {
let a = i64::try_from(r.bytes).unwrap_or(i64::MAX);
let b = i64::try_from(r.bytes_b.unwrap_or(0)).unwrap_or(i64::MAX);
(
r.tensors.to_string(),
format_size(r.bytes),
format_size(r.bytes_b.unwrap_or(0)),
format_size_delta(b.saturating_sub(a)),
)
})
.collect();
let pw = column_width("Pattern", rows.iter().map(|r| r.pattern.as_str()));
let tw = column_width("Tensors", cells.iter().map(|(t, ..)| t.as_str()));
let aw = column_width("A Bytes", cells.iter().map(|(_, a, ..)| a.as_str()));
let bw = column_width("B Bytes", cells.iter().map(|(_, _, b, _)| b.as_str()));
let dw = column_width("\u{0394} Bytes", cells.iter().map(|(.., d)| d.as_str()));
println!(
" {:<pw$} {:>tw$} {:>aw$} {:>bw$} {:>dw$}",
"Pattern", "Tensors", "A Bytes", "B Bytes", "\u{0394} Bytes",
);
for (r, (tensors, a_bytes, b_bytes, delta)) in rows.iter().zip(&cells).take(cap) {
println!(
" {:<pw$} {tensors:>tw$} {a_bytes:>aw$} {b_bytes:>bw$} {delta:>dw$}",
r.pattern
);
}
print_truncation_note(cap, rows.len());
println!();
}
#[allow(clippy::too_many_arguments)]
fn print_diff_collapse(
repo_a: &str,
repo_b: &str,
tensors_a: &HashMap<String, inspect::TensorInfo>,
tensors_b: &HashMap<String, inspect::TensorInfo>,
only_a: &[&str],
only_b: &[&str],
differ: &[&str],
filter: Option<&str>,
limit: Option<usize>,
) {
print_diff_repo_header(repo_a, repo_b);
println!();
let cap = limit.unwrap_or(usize::MAX);
let rows_a = aggregate_diff_collapse(only_a, tensors_a, None);
if !rows_a.is_empty() {
print_diff_collapse_section("Only in A", &rows_a, cap, only_a.len());
}
let rows_b = aggregate_diff_collapse(only_b, tensors_b, None);
if !rows_b.is_empty() {
print_diff_collapse_section("Only in B", &rows_b, cap, only_b.len());
}
let rows_differ = aggregate_diff_collapse(differ, tensors_a, Some(tensors_b));
if !rows_differ.is_empty() {
print_diff_collapse_differ(&rows_differ, cap, differ.len());
}
println!(" {}", "\u{2500}".repeat(70));
print!(
" only-A: {} {} ({} {}) | only-B: {} {} ({} {}) | differ: {} {} ({} {})",
only_a.len(),
pluralize(only_a.len(), "tensor", "tensors"),
rows_a.len(),
pluralize(rows_a.len(), "pattern", "patterns"),
only_b.len(),
pluralize(only_b.len(), "tensor", "tensors"),
rows_b.len(),
pluralize(rows_b.len(), "pattern", "patterns"),
differ.len(),
pluralize(differ.len(), "tensor", "tensors"),
rows_differ.len(),
pluralize(rows_differ.len(), "pattern", "patterns"),
);
if let Some(pattern) = filter {
println!(" (filter: {pattern:?})");
} else {
println!();
}
}
#[derive(serde::Serialize)]
struct ConfigFieldDiff {
field: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
a: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
b: Option<String>,
differs: bool,
}
fn opt_to_string<T: std::fmt::Display>(v: Option<&T>) -> Option<String> {
v.map(std::string::ToString::to_string)
}
fn opt_vec_to_string<T: std::fmt::Display>(v: Option<&[T]>) -> Option<String> {
v.map(|items| {
items
.iter()
.map(std::string::ToString::to_string)
.collect::<Vec<_>>()
.join(", ")
})
}
const CONFIG_CELL_DISPLAY_CAP: usize = 60;
fn truncate_for_display(s: &str) -> std::borrow::Cow<'_, str> {
if s.chars().count() <= CONFIG_CELL_DISPLAY_CAP {
std::borrow::Cow::Borrowed(s)
} else {
let head: String = s.chars().take(CONFIG_CELL_DISPLAY_CAP).collect();
std::borrow::Cow::Owned(format!("{head}\u{2026}"))
}
}
fn build_config_diff_rows(
a: &inspect::ModelConfig,
b: &inspect::ModelConfig,
) -> Vec<ConfigFieldDiff> {
let row = |field: &'static str, a_val: Option<String>, b_val: Option<String>| {
let differs = a_val != b_val;
ConfigFieldDiff {
field,
a: a_val,
b: b_val,
differs,
}
};
vec![
row("model_type", a.model_type.clone(), b.model_type.clone()),
row(
"num_hidden_layers",
opt_to_string(a.num_hidden_layers.as_ref()),
opt_to_string(b.num_hidden_layers.as_ref()),
),
row(
"num_attention_heads",
opt_to_string(a.num_attention_heads.as_ref()),
opt_to_string(b.num_attention_heads.as_ref()),
),
row(
"num_key_value_heads",
opt_to_string(a.num_key_value_heads.as_ref()),
opt_to_string(b.num_key_value_heads.as_ref()),
),
row(
"head_dim",
opt_to_string(a.head_dim.as_ref()),
opt_to_string(b.head_dim.as_ref()),
),
row(
"hidden_size",
opt_to_string(a.hidden_size.as_ref()),
opt_to_string(b.hidden_size.as_ref()),
),
row("torch_dtype", a.torch_dtype.clone(), b.torch_dtype.clone()),
row(
"sliding_window",
opt_to_string(a.sliding_window.as_ref()),
opt_to_string(b.sliding_window.as_ref()),
),
row(
"sliding_window_pattern",
opt_to_string(a.sliding_window_pattern.as_ref()),
opt_to_string(b.sliding_window_pattern.as_ref()),
),
row(
"use_sliding_window",
opt_to_string(a.use_sliding_window.as_ref()),
opt_to_string(b.use_sliding_window.as_ref()),
),
row(
"kv_lora_rank",
opt_to_string(a.kv_lora_rank.as_ref()),
opt_to_string(b.kv_lora_rank.as_ref()),
),
row(
"qk_rope_head_dim",
opt_to_string(a.qk_rope_head_dim.as_ref()),
opt_to_string(b.qk_rope_head_dim.as_ref()),
),
row(
"layer_types",
opt_vec_to_string(a.layer_types.as_deref()),
opt_vec_to_string(b.layer_types.as_deref()),
),
row(
"hybrid_override_pattern",
a.hybrid_override_pattern.clone(),
b.hybrid_override_pattern.clone(),
),
row(
"attn_layer_indices",
opt_vec_to_string(a.attn_layer_indices.as_deref()),
opt_vec_to_string(b.attn_layer_indices.as_deref()),
),
row(
"full_attention_interval",
opt_to_string(a.full_attention_interval.as_ref()),
opt_to_string(b.full_attention_interval.as_ref()),
),
row(
"attn_layer_offset",
opt_to_string(a.attn_layer_offset.as_ref()),
opt_to_string(b.attn_layer_offset.as_ref()),
),
row(
"attn_layer_period",
opt_to_string(a.attn_layer_period.as_ref()),
opt_to_string(b.attn_layer_period.as_ref()),
),
row(
"mamba_n_heads",
opt_to_string(a.mamba_n_heads.as_ref()),
opt_to_string(b.mamba_n_heads.as_ref()),
),
row(
"mamba_d_head",
opt_to_string(a.mamba_d_head.as_ref()),
opt_to_string(b.mamba_d_head.as_ref()),
),
row(
"mamba_d_state",
opt_to_string(a.mamba_d_state.as_ref()),
opt_to_string(b.mamba_d_state.as_ref()),
),
row(
"mamba_d_conv",
opt_to_string(a.mamba_d_conv.as_ref()),
opt_to_string(b.mamba_d_conv.as_ref()),
),
row(
"mamba_n_groups",
opt_to_string(a.mamba_n_groups.as_ref()),
opt_to_string(b.mamba_n_groups.as_ref()),
),
]
}
#[derive(serde::Serialize)]
struct DiffConfigResult {
repo_a: String,
repo_b: String,
#[serde(skip_serializing_if = "Vec::is_empty")]
missing: Vec<String>,
fields: Vec<ConfigFieldDiff>,
differing_count: usize,
}
fn missing_config_hint(cached: bool) -> &'static str {
if cached {
"Hint: not every repo ships a config.json (e.g. non-model repos) — or it just isn't cached locally yet; retry without --cached."
} else {
"Hint: not every repo ships a config.json (e.g. non-model repos)."
}
}
#[allow(clippy::too_many_arguments)] fn run_diff_config(
repo_a: &str,
repo_b: &str,
revision_a: Option<&str>,
revision_b: Option<&str>,
token: Option<&str>,
cached: bool,
all: bool,
json: bool,
) -> Result<(), FetchError> {
let (config_a, config_b) = if cached {
(
inspect::fetch_model_config_cached(repo_a, revision_a)?,
inspect::fetch_model_config_cached(repo_b, revision_b)?,
)
} else {
let resolved_token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let (result_a, result_b) = rt.block_on(async {
tokio::join!(
inspect::fetch_model_config(repo_a, resolved_token.as_deref(), revision_a),
inspect::fetch_model_config(repo_b, resolved_token.as_deref(), revision_b),
)
});
(
result_a.map_err(|e| enrich_gated_content_error(e, repo_a, token))?,
result_b.map_err(|e| enrich_gated_content_error(e, repo_b, token))?,
)
};
let missing: Vec<&str> = [(repo_a, &config_a), (repo_b, &config_b)]
.into_iter()
.filter_map(|(repo, config)| config.is_none().then_some(repo))
.collect();
if !missing.is_empty() {
if json {
let result = DiffConfigResult {
repo_a: repo_a.to_owned(),
repo_b: repo_b.to_owned(),
missing: missing.into_iter().map(str::to_owned).collect(),
fields: Vec::new(),
differing_count: 0,
};
return emit_json(&result);
}
for repo in missing {
println!("No config.json found in {repo}.");
println!("{}", missing_config_hint(cached));
}
return Ok(());
}
let (Some(config_a), Some(config_b)) = (config_a, config_b) else {
unreachable!("missing.is_empty() guarantees both configs are Some")
};
let fields = build_config_diff_rows(&config_a, &config_b);
let differing_count = fields.iter().filter(|r| r.differs).count();
if json {
let result = DiffConfigResult {
repo_a: repo_a.to_owned(),
repo_b: repo_b.to_owned(),
missing: Vec::new(),
fields,
differing_count,
};
return emit_json(&result);
}
print_diff_repo_header(repo_a, repo_b);
println!();
let shown: Vec<&ConfigFieldDiff> = if all {
fields.iter().collect()
} else {
fields.iter().filter(|r| r.differs).collect()
};
let cells: Vec<(std::borrow::Cow<'_, str>, std::borrow::Cow<'_, str>)> = shown
.iter()
.map(|r| {
(
truncate_for_display(r.a.as_deref().unwrap_or("\u{2014}")),
truncate_for_display(r.b.as_deref().unwrap_or("\u{2014}")),
)
})
.collect();
let fw = column_width("Field", shown.iter().map(|r| r.field));
let aw = column_width("A", cells.iter().map(|(a, _)| a.as_ref()));
let bw = column_width("B", cells.iter().map(|(_, b)| b.as_ref()));
println!(" {:<fw$} {:<aw$} {:<bw$}", "Field", "A", "B");
for (r, (a, b)) in shown.iter().zip(&cells) {
println!(" {:<fw$} {a:<aw$} {b:<bw$}", r.field);
}
println!();
println!(" {differing_count} of {} fields differ", fields.len());
Ok(())
}
#[allow(clippy::fn_params_excessive_bools, clippy::too_many_arguments)]
fn run_inspect(
repo_id: &str,
filename: Option<&str>,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
cache_headers: bool,
list: bool,
pick: bool,
no_metadata: bool,
json: bool,
filter: Option<&str>,
dtypes: bool,
group_by: Option<&str>,
limit: Option<usize>,
tree: bool,
check_gpu: Option<u32>,
context: Option<u32>,
) -> Result<(), FetchError> {
if list {
return run_inspect_list(repo_id, revision, token, cached);
}
if pick {
let picked = pick_inspect_file(repo_id, filename, revision, token, cached)?;
return run_inspect_single(
repo_id,
picked.as_str(),
revision,
token,
cached,
cache_headers,
no_metadata,
json,
filter,
dtypes,
group_by,
limit,
tree,
check_gpu,
context,
)
.map_err(|e| enrich_gated_content_error(e, repo_id, token));
}
match filename {
Some(f) => {
let resolved = resolve_inspect_filename_arg(f, repo_id, revision, token, cached)?;
run_inspect_single(
repo_id,
resolved.as_str(),
revision,
token,
cached,
cache_headers,
no_metadata,
json,
filter,
dtypes,
group_by,
limit,
tree,
check_gpu,
context,
)
.map_err(|e| enrich_gated_content_error(e, repo_id, token))
}
None => run_inspect_repo(
repo_id, revision, token, cached, json, filter, dtypes, group_by, limit, tree,
check_gpu, context,
)
.map_err(|e| enrich_gated_content_error(e, repo_id, token)),
}
}
#[allow(clippy::fn_params_excessive_bools, clippy::too_many_arguments)]
fn run_peek(
repo_id: &str,
filename: &str,
revision: Option<&str>,
token: Option<&str>,
head: Option<u64>,
tail: Option<u64>,
bytes: bool,
gunzip: bool,
no_gunzip: bool,
max: u64,
) -> Result<(), FetchError> {
let mode = peek::resolve_mode(head, tail, bytes)?;
let effective_gunzip = peek::resolve_gunzip(gunzip, no_gunzip, filename);
let options = peek::PeekOptions::new(mode, effective_gunzip, max, filename.to_owned());
let owned_token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let outcome = rt
.block_on(peek::peek(
repo_id,
filename,
owned_token.as_deref(),
revision,
options,
))
.map_err(|e| enrich_gated_content_error(e, repo_id, owned_token.as_deref()))?;
std::io::stdout()
.write_all(&outcome.content)
.map_err(|e| FetchError::Io {
path: PathBuf::from("<stdout>"),
source: e,
})?;
if let Some(note) = outcome.truncated {
eprintln!("{note}");
}
Ok(())
}
fn is_auth_status_error(err: &FetchError) -> bool {
matches!(err, FetchError::Http(msg)
if msg.contains("returned status 401") || msg.contains("returned status 403"))
}
fn enrich_gated_content_error(err: FetchError, repo_id: &str, token: Option<&str>) -> FetchError {
if !is_auth_status_error(&err) {
return err;
}
let Ok(rt) = tokio::runtime::Runtime::new() else {
return err;
};
let Ok(metadata) = rt.block_on(discover::fetch_model_card(repo_id)) else {
return err;
};
if !metadata.gated.is_gated() {
return err;
}
let effective_token = token
.map(ToOwned::to_owned)
.or_else(|| std::env::var("HF_TOKEN").ok());
let reason = if effective_token.is_none() {
format!(
"{repo_id} is a gated model — its file listing is public but content \
requires access: accept the license at https://huggingface.co/{repo_id} \
and set HF_TOKEN or pass --token"
)
} else {
format!(
"{repo_id} is a gated model and your token was rejected — accept the \
license at https://huggingface.co/{repo_id} (each gated family is \
licensed separately) and check that the token grants gated-repo read access"
)
};
FetchError::Auth { reason }
}
fn resolve_inspect_filename_arg(
arg: &str,
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Result<String, FetchError> {
let Ok(n) = arg.parse::<usize>() else {
return Ok(arg.to_owned());
};
let (entries, commit_sha) = gather_tensor_listing(repo_id, revision, token, cached)?;
if entries.is_empty() {
return Err(FetchError::InvalidArgument(format!(
"index {n} cannot be resolved: no supported tensor files \
(.safetensors / .gguf / .npz / .pth) in repository {repo_id} \
(run `hf-fm inspect {repo_id} --list` to confirm)"
)));
}
if n == 0 || n > entries.len() {
return Err(FetchError::InvalidArgument(format!(
"index {n} is out of range (repository has {count} tensor files — \
use 1..{count}; run `hf-fm inspect {repo_id} --list` to see them)",
count = entries.len()
)));
}
#[allow(clippy::indexing_slicing)]
let (filename, _size) = &entries[n - 1];
let rev_note = match &commit_sha {
Some(sha) => format!(" (repo rev: {})", short_sha(sha)),
None => String::new(),
};
eprintln!("Resolving index {n} → {filename}{rev_note}");
Ok(filename.clone())
}
fn short_sha(sha: &str) -> String {
sha.get(..12).map_or_else(|| sha.to_owned(), str::to_owned)
}
fn gather_tensor_listing(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Result<inspect::TensorFileListing, FetchError> {
if cached {
return inspect::list_cached_tensor_files(repo_id, revision);
}
let resolved_token = token
.map(ToOwned::to_owned)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let client = hf_fetch_model::build_client(resolved_token.as_deref())?;
let (files, commit_sha) = rt.block_on(repo::list_repo_files_with_commit(
repo_id,
resolved_token.as_deref(),
revision,
&client,
))?;
let mut entries: Vec<(String, u64)> = files
.into_iter()
.filter(|f| inspect::is_supported_tensor_file(f.filename.as_str()))
.map(|f| (f.filename, f.size.unwrap_or(0)))
.collect();
entries.sort_by(|a, b| a.0.cmp(&b.0));
Ok((entries, commit_sha))
}
fn pick_inspect_file(
repo_id: &str,
needle: Option<&str>,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Result<String, FetchError> {
if !(std::io::stdin().is_terminal() && std::io::stderr().is_terminal()) {
return Err(FetchError::InvalidArgument(
"--pick requires an interactive terminal (stdin + stderr attached): \
run `hf-fm inspect <repo> --list`, then `hf-fm inspect <repo> <n>`"
.to_owned(),
));
}
let (entries, commit_sha) = gather_tensor_listing(repo_id, revision, token, cached)?;
let candidates = narrow_pick_candidates(&entries, needle);
let rev_note = match &commit_sha {
Some(sha) => format!(" (repo rev: {})", short_sha(sha)),
None => String::new(),
};
match candidates.as_slice() {
[] => Err(FetchError::InvalidArgument(match needle {
Some(n) => format!(
"no tensor files match {n:?} in {repo_id} \
(run `hf-fm inspect {repo_id} --list` to see all)"
),
None => format!(
"no supported tensor files (.safetensors / .gguf / .npz / .pth) \
in repository {repo_id}"
),
})),
[(filename, _size)] => {
eprintln!("Resolving to {filename}{rev_note}");
Ok((*filename).clone())
}
_ => prompt_pick_selection(repo_id, needle, &candidates, &rev_note),
}
}
fn prompt_pick_selection(
repo_id: &str,
needle: Option<&str>,
candidates: &[&(String, u64)],
rev_note: &str,
) -> Result<String, FetchError> {
match needle {
Some(n) => eprintln!("Multiple tensor files match {n:?} in {repo_id}:"),
None => eprintln!("Tensor files in {repo_id}:"),
}
let count = candidates.len();
let index_width = count.to_string().len();
let file_width = candidates.iter().map(|(f, _)| f.len()).max().unwrap_or(0);
let size_strings: Vec<String> = candidates.iter().map(|(_, s)| format_size(*s)).collect();
let size_width = size_strings.iter().map(String::len).max().unwrap_or(0);
for (i, ((filename, _), size_str)) in candidates.iter().zip(size_strings.iter()).enumerate() {
let n = i + 1;
eprintln!(" {n:>index_width$} {filename:<file_width$} {size_str:>size_width$}");
}
loop {
eprint!("Pick [1..{count}]: ");
let mut input = String::new();
let bytes_read = std::io::stdin()
.read_line(&mut input)
.map_err(|e| FetchError::Io {
path: PathBuf::from("<stdin>"),
source: e,
})?;
let trimmed = input.trim();
if bytes_read == 0 || trimmed.is_empty() {
return Err(FetchError::InvalidArgument(
"cancelled — no file picked".to_owned(),
));
}
if let Some(n) = parse_pick_input(trimmed, count) {
#[allow(clippy::indexing_slicing)]
let (filename, _size) = candidates[n - 1];
eprintln!("Resolving to {filename}{rev_note}");
return Ok(filename.clone());
}
eprintln!(
"invalid choice {trimmed:?} — enter a number between 1 and {count}, \
or press Enter to cancel"
);
}
}
fn narrow_pick_candidates<'a>(
entries: &'a [(String, u64)],
needle: Option<&str>,
) -> Vec<&'a (String, u64)> {
match needle {
None => entries.iter().collect(),
Some(raw) => {
let needle_lc = raw.to_lowercase();
entries
.iter()
.filter(|(filename, _)| filename.to_lowercase().contains(&needle_lc))
.collect()
}
}
}
fn parse_pick_input(line: &str, count: usize) -> Option<usize> {
line.trim()
.parse::<usize>()
.ok()
.filter(|n| (1..=count).contains(n))
}
fn run_inspect_list(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Result<(), FetchError> {
let (entries, commit_sha) = gather_tensor_listing(repo_id, revision, token, cached)?;
println!("Repo: {repo_id}");
let rev_label = revision.unwrap_or("main");
match &commit_sha {
Some(sha) => println!("Rev: {sha} ({rev_label})"),
None => println!("Rev: (unknown) ({rev_label})"),
}
println!();
if entries.is_empty() {
println!(
"No supported tensor files (.safetensors / .gguf / .npz / .pth) in this repository."
);
if cached {
println!();
println!("Hint: the repo may not be cached locally. Try without --cached.");
}
return Ok(());
}
let count = entries.len();
let index_width = count.to_string().len();
let file_width = entries
.iter()
.map(|(f, _)| f.len())
.max()
.unwrap_or(4)
.max(4); let size_strings: Vec<String> = entries.iter().map(|(_, s)| format_size(*s)).collect();
let size_width = size_strings
.iter()
.map(String::len)
.max()
.unwrap_or(4)
.max(4);
println!(
"{:>index_width$} {:<file_width$} {:>size_width$}",
"#", "File", "Size"
);
println!(
"{:->index_width$} {:-<file_width$} {:->size_width$}",
"", "", ""
);
let mut total: u64 = 0;
for (i, ((filename, size), size_str)) in entries.iter().zip(size_strings.iter()).enumerate() {
let n = i + 1;
println!("{n:>index_width$} {filename:<file_width$} {size_str:>size_width$}");
total = total.saturating_add(*size);
}
println!();
println!(
"{count} {}, {} total",
pluralize(count, "file", "files"),
format_size(total)
);
if revision.is_none() {
if let Some(sha) = commit_sha.as_deref() {
println!();
println!(
"Tip: run `hf-fm inspect {repo_id} <n>` to inspect file #n.\n \
Pass `--revision {sha}` on both sides to lock against this view."
);
} else {
println!();
println!("Tip: run `hf-fm inspect {repo_id} <n>` to inspect file #n.");
}
}
Ok(())
}
#[allow(clippy::exhaustive_enums)] #[derive(Debug, Clone)]
enum TreeNode {
Leaf(LeafNode),
Branch(BranchNode),
Ranged(RangedNode),
}
#[derive(Debug, Clone)]
struct LeafNode {
name: String,
dtype: String,
shape: Vec<usize>,
params: u64,
bytes: u64,
}
#[derive(Debug, Clone)]
struct BranchNode {
segment: String,
children: Vec<TreeNode>,
total_tensors: usize,
total_params: u64,
total_bytes: u64,
}
#[derive(Debug, Clone)]
struct RangedNode {
segment: String,
range_start: usize,
range_end: usize,
template: Vec<TreeNode>,
total_tensors: usize,
total_params: u64,
total_bytes: u64,
}
impl TreeNode {
fn total_tensors(&self) -> usize {
match self {
Self::Leaf(_) => 1,
Self::Branch(b) => b.total_tensors,
Self::Ranged(r) => r.total_tensors,
}
}
fn total_params(&self) -> u64 {
match self {
Self::Leaf(l) => l.params,
Self::Branch(b) => b.total_params,
Self::Ranged(r) => r.total_params,
}
}
fn total_bytes(&self) -> u64 {
match self {
Self::Leaf(l) => l.bytes,
Self::Branch(b) => b.total_bytes,
Self::Ranged(r) => r.total_bytes,
}
}
}
#[derive(Debug, Default)]
struct TrieNode {
tensor: Option<inspect::TensorInfo>,
children: BTreeMap<String, TrieNode>,
}
fn build_tree(tensors: &[inspect::TensorInfo]) -> Vec<TreeNode> {
let mut root = TrieNode::default();
for t in tensors {
let segments: Vec<&str> = t.name.as_str().split('.').collect();
insert_trie(&mut root, &segments, t.clone());
}
root.children
.into_iter()
.map(|(seg, child)| trie_to_tree(seg, child))
.collect()
}
fn insert_trie(node: &mut TrieNode, segments: &[&str], tensor: inspect::TensorInfo) {
let Some((head, rest)) = segments.split_first() else {
node.tensor = Some(tensor);
return;
};
if rest.is_empty() {
let child = node.children.entry((*head).to_owned()).or_default();
child.tensor = Some(tensor);
} else {
let child = node.children.entry((*head).to_owned()).or_default();
insert_trie(child, rest, tensor);
}
}
fn trie_to_tree(segment: String, mut node: TrieNode) -> TreeNode {
if node.children.is_empty() {
if let Some(tensor) = node.tensor {
let params = tensor.num_elements();
let bytes = tensor.byte_len();
return TreeNode::Leaf(LeafNode {
name: segment,
dtype: tensor.dtype,
shape: tensor.shape,
params,
bytes,
});
}
return TreeNode::Branch(BranchNode {
segment,
children: Vec::new(),
total_tensors: 0,
total_params: 0,
total_bytes: 0,
});
}
if node.tensor.is_none()
&& node.children.len() == 1
&& let Some((child_segment, child)) = node.children.pop_first()
{
let merged = format!("{segment}.{child_segment}");
return trie_to_tree(merged, child);
}
let mut entries: Vec<(String, TrieNode)> = node.children.into_iter().collect();
if entries
.iter()
.all(|(seg, _)| seg.as_str().parse::<usize>().is_ok())
{
entries.sort_by_key(|(seg, _)| {
seg.as_str().parse::<usize>().unwrap_or(usize::MAX)
});
}
let mut children: Vec<TreeNode> = entries
.into_iter()
.map(|(seg, child)| trie_to_tree(seg, child))
.collect();
if let Some(tensor) = node.tensor {
let params = tensor.num_elements();
let bytes = tensor.byte_len();
children.insert(
0,
TreeNode::Leaf(LeafNode {
name: String::new(),
dtype: tensor.dtype,
shape: tensor.shape,
params,
bytes,
}),
);
}
let total_tensors: usize = children.iter().map(TreeNode::total_tensors).sum();
let total_params: u64 = children
.iter()
.map(TreeNode::total_params)
.fold(0u64, u64::saturating_add);
let total_bytes: u64 = children
.iter()
.map(TreeNode::total_bytes)
.fold(0u64, u64::saturating_add);
TreeNode::Branch(BranchNode {
segment,
children,
total_tensors,
total_params,
total_bytes,
})
}
fn collapse_ranges(nodes: Vec<TreeNode>) -> Vec<TreeNode> {
nodes.into_iter().map(collapse_node).collect()
}
fn collapse_node(node: TreeNode) -> TreeNode {
match node {
TreeNode::Leaf(_) => node,
TreeNode::Branch(mut branch) => {
branch.children = collapse_ranges(branch.children);
try_collapse_range(&branch).unwrap_or(TreeNode::Branch(branch))
}
TreeNode::Ranged(mut ranged) => {
ranged.template = collapse_ranges(ranged.template);
TreeNode::Ranged(ranged)
}
}
}
const MAX_EDGE_OUTLIERS: usize = 3;
fn try_collapse_range(branch: &BranchNode) -> Option<TreeNode> {
if branch.children.len() < 2 {
return None;
}
let mut indexed: Vec<(usize, &BranchNode)> = Vec::with_capacity(branch.children.len());
for child in &branch.children {
let TreeNode::Branch(sub) = child else {
return None;
};
let idx: usize = sub.segment.as_str().parse().ok()?;
indexed.push((idx, sub));
}
indexed.sort_by_key(|(i, _)| *i);
for (expected, (actual, _)) in indexed.iter().enumerate() {
if expected != *actual {
return None;
}
}
let n = indexed.len();
let mut candidates: Vec<(usize, usize)> = (0..=MAX_EDGE_OUTLIERS)
.flat_map(|skip_front| {
(0..=MAX_EDGE_OUTLIERS).map(move |skip_back| (skip_front, skip_back))
})
.collect();
candidates.sort_by_key(|&(skip_front, skip_back)| (skip_front + skip_back, skip_back));
for (skip_front, skip_back) in candidates {
if skip_front + skip_back >= n {
continue;
}
#[allow(clippy::indexing_slicing)]
let inner = &indexed[skip_front..n - skip_back];
if inner.len() < 2 {
continue;
}
#[allow(clippy::indexing_slicing)]
let (start_idx, template_branch) = inner[0];
#[allow(clippy::indexing_slicing)]
let uniform = inner[1..]
.iter()
.all(|(_, other)| branches_structurally_equal(template_branch, other));
if !uniform {
continue;
}
let count = inner.len();
#[allow(clippy::as_conversions)]
let count_u64 = count as u64;
let ranged = RangedNode {
segment: if skip_front == 0 && skip_back == 0 {
branch.segment.clone()
} else {
String::new()
},
range_start: start_idx,
range_end: start_idx.saturating_add(count).saturating_sub(1),
template: template_branch.children.clone(),
total_tensors: template_branch.total_tensors.saturating_mul(count),
total_params: template_branch.total_params.saturating_mul(count_u64),
total_bytes: template_branch.total_bytes.saturating_mul(count_u64),
};
if skip_front == 0 && skip_back == 0 {
return Some(TreeNode::Ranged(ranged));
}
let mut children =
Vec::with_capacity(skip_front.saturating_add(1).saturating_add(skip_back));
#[allow(clippy::indexing_slicing)]
for (_, b) in &indexed[..skip_front] {
children.push(TreeNode::Branch((*b).clone()));
}
children.push(TreeNode::Ranged(ranged));
#[allow(clippy::indexing_slicing)]
for (_, b) in &indexed[n - skip_back..] {
children.push(TreeNode::Branch((*b).clone()));
}
return Some(TreeNode::Branch(BranchNode {
segment: branch.segment.clone(),
children,
total_tensors: branch.total_tensors,
total_params: branch.total_params,
total_bytes: branch.total_bytes,
}));
}
None
}
fn branches_structurally_equal(a: &BranchNode, b: &BranchNode) -> bool {
if a.children.len() != b.children.len() {
return false;
}
a.children
.iter()
.zip(b.children.iter())
.all(|(c1, c2)| nodes_structurally_equal(c1, c2))
}
fn nodes_structurally_equal(a: &TreeNode, b: &TreeNode) -> bool {
match (a, b) {
(TreeNode::Leaf(l1), TreeNode::Leaf(l2)) => {
l1.name == l2.name && l1.dtype == l2.dtype && l1.shape == l2.shape
}
(TreeNode::Branch(b1), TreeNode::Branch(b2)) => {
b1.segment == b2.segment && branches_structurally_equal(b1, b2)
}
(TreeNode::Ranged(r1), TreeNode::Ranged(r2)) => {
r1.segment == r2.segment
&& r1.range_end == r2.range_end
&& r1.template.len() == r2.template.len()
&& r1
.template
.iter()
.zip(r2.template.iter())
.all(|(c1, c2)| nodes_structurally_equal(c1, c2))
}
_ => false,
}
}
fn render_tree(nodes: &[TreeNode]) {
render_children(nodes, "");
}
fn render_children(children: &[TreeNode], prefix: &str) {
let leaf_name_width: usize = children
.iter()
.filter_map(|c| match c {
TreeNode::Leaf(l) => Some(l.name.len()),
TreeNode::Branch(_) | TreeNode::Ranged(_) => None,
})
.max()
.unwrap_or(0);
let dtype_width: usize = children
.iter()
.filter_map(|c| match c {
TreeNode::Leaf(l) => Some(l.dtype.len()),
TreeNode::Branch(_) | TreeNode::Ranged(_) => None,
})
.max()
.unwrap_or(0);
for (i, child) in children.iter().enumerate() {
let is_last = i + 1 == children.len();
render_node(child, prefix, is_last, leaf_name_width, dtype_width);
}
}
fn render_node(
node: &TreeNode,
prefix: &str,
is_last: bool,
leaf_name_width: usize,
dtype_width: usize,
) {
let connector = if is_last { "└── " } else { "├── " };
let indent = if is_last { " " } else { "│ " };
match node {
TreeNode::Leaf(leaf) => {
let shape_str = format!("{:?}", leaf.shape);
let size_str = format_size(leaf.bytes);
println!(
" {prefix}{connector}{name:<nw$} {dtype:<dw$} {shape_str} {size_str}",
name = leaf.name,
dtype = leaf.dtype,
nw = leaf_name_width,
dw = dtype_width,
);
}
TreeNode::Branch(branch) => {
println!(
" {prefix}{connector}{seg}.",
seg = branch.segment.as_str(), );
let new_prefix = format!("{prefix}{indent}");
render_children(&branch.children, new_prefix.as_str()); }
TreeNode::Ranged(ranged) => {
let count = ranged.range_end - ranged.range_start + 1;
let range_label = format!(
"[{start}..{end}]",
start = ranged.range_start,
end = ranged.range_end
);
if ranged.segment.is_empty() {
println!(" {prefix}{connector}{range_label}. (\u{00d7}{count})");
} else {
println!(
" {prefix}{connector}{seg}.{range_label}. (\u{00d7}{count})",
seg = ranged.segment.as_str(), );
}
let new_prefix = format!("{prefix}{indent}");
render_children(&ranged.template, new_prefix.as_str()); }
}
}
#[derive(serde::Serialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum TreeJsonNode<'a> {
Leaf {
name: &'a str,
dtype: &'a str,
shape: &'a [usize],
params: u64,
bytes: u64,
},
Branch {
name: &'a str,
tensors: usize,
params: u64,
bytes: u64,
children: Vec<TreeJsonNode<'a>>,
},
Ranged {
name: &'a str,
range_start: usize,
range_end: usize,
count: usize,
tensors: usize,
params: u64,
bytes: u64,
template: Vec<TreeJsonNode<'a>>,
},
}
#[derive(serde::Serialize)]
struct TreeJsonOutput<'a> {
repo_id: &'a str,
filename: &'a str,
total_tensors: usize,
total_params: u64,
tree: Vec<TreeJsonNode<'a>>,
#[serde(skip_serializing_if = "Option::is_none")]
gpu_check: Option<serde_json::Value>,
}
fn tree_to_json(nodes: &[TreeNode]) -> Vec<TreeJsonNode<'_>> {
nodes.iter().map(node_to_json).collect()
}
fn print_tree_summary(
tensors: &[inspect::TensorInfo],
filter: Option<&str>,
total_tensor_count: usize,
total_params: u64,
) {
let forest = collapse_ranges(build_tree(tensors));
println!();
render_tree(&forest);
let shown: usize = forest.iter().map(TreeNode::total_tensors).sum();
let shown_params: u64 = forest
.iter()
.map(TreeNode::total_params)
.fold(0u64, u64::saturating_add);
let tensor_label = if shown == 1 { "tensor" } else { "tensors" };
if let Some(pattern) = filter {
println!(
" {shown}/{total_tensor_count} {tensor_label}, {}/{} params (filter: {pattern:?})",
inspect::format_params(shown_params),
inspect::format_params(total_params),
);
} else {
println!(
" {shown} {tensor_label}, {} params",
inspect::format_params(shown_params),
);
}
}
fn print_tree_json(
repo_id: &str,
filename: &str,
tensors: &[inspect::TensorInfo],
total_tensor_count: usize,
total_params: u64,
gpu_check: Option<serde_json::Value>,
) -> Result<(), FetchError> {
let forest = collapse_ranges(build_tree(tensors));
let output = TreeJsonOutput {
repo_id,
filename,
total_tensors: total_tensor_count,
total_params,
tree: tree_to_json(&forest),
gpu_check,
};
let serialized = serde_json::to_string_pretty(&output)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{serialized}");
Ok(())
}
fn node_to_json(node: &TreeNode) -> TreeJsonNode<'_> {
match node {
TreeNode::Leaf(l) => TreeJsonNode::Leaf {
name: l.name.as_str(), dtype: l.dtype.as_str(), shape: l.shape.as_slice(), params: l.params,
bytes: l.bytes,
},
TreeNode::Branch(b) => TreeJsonNode::Branch {
name: b.segment.as_str(), tensors: b.total_tensors,
params: b.total_params,
bytes: b.total_bytes,
children: tree_to_json(&b.children),
},
TreeNode::Ranged(r) => {
let count = r.range_end - r.range_start + 1;
TreeJsonNode::Ranged {
name: r.segment.as_str(), range_start: r.range_start,
range_end: r.range_end,
count,
tensors: r.total_tensors,
params: r.total_params,
bytes: r.total_bytes,
template: tree_to_json(&r.template),
}
}
}
}
#[derive(serde::Serialize)]
struct TruncationInfo {
shown: usize,
total: usize,
}
#[derive(serde::Serialize)]
struct InspectJsonOutput<'a> {
#[serde(flatten)]
header: &'a inspect::SafetensorsHeaderInfo,
#[serde(skip_serializing_if = "Option::is_none")]
truncated: Option<TruncationInfo>,
#[serde(skip_serializing_if = "Option::is_none")]
gpu_check: Option<serde_json::Value>,
}
async fn dispatch_inspect_remote(
repo_id: &str,
filename: &str,
revision: Option<&str>,
token: Option<&str>,
is_npz: bool,
is_gguf: bool,
is_pth: bool,
) -> Result<
(
inspect::SafetensorsHeaderInfo,
inspect::InspectSource,
Option<RangeStats>,
),
FetchError,
> {
if is_npz {
inspect::inspect_npz(repo_id, filename, token, revision).await
} else if is_gguf {
inspect::inspect_gguf(repo_id, filename, token, revision).await
} else if is_pth {
inspect::inspect_pth(repo_id, filename, token, revision).await
} else {
inspect::inspect_safetensors(repo_id, filename, token, revision).await
}
}
async fn dispatch_inspect_remote_from_reader(
reader: HttpRangeReader,
filename: &str,
is_npz: bool,
is_gguf: bool,
is_pth: bool,
) -> Result<
(
inspect::SafetensorsHeaderInfo,
inspect::InspectSource,
Option<RangeStats>,
),
FetchError,
> {
if is_npz {
inspect::inspect_npz_from_reader(reader, filename).await
} else if is_gguf {
inspect::inspect_gguf_from_reader(reader, filename).await
} else if is_pth {
inspect::inspect_pth_from_reader(reader, filename).await
} else {
inspect::inspect_safetensors_from_reader(reader, filename).await
}
}
#[allow(clippy::fn_params_excessive_bools, clippy::too_many_arguments)]
async fn inspect_remote_with_cache(
repo_id: &str,
filename: &str,
revision: Option<&str>,
token: Option<&str>,
is_npz: bool,
is_gguf: bool,
is_pth: bool,
cache_headers: bool,
) -> Result<
(
inspect::SafetensorsHeaderInfo,
inspect::InspectSource,
Option<RangeStats>,
),
FetchError,
> {
let rev = revision.unwrap_or("main");
if !cache_headers || inspect::resolve_cached_path(repo_id, rev, filename).is_some() {
return dispatch_inspect_remote(
repo_id, filename, revision, token, is_npz, is_gguf, is_pth,
)
.await;
}
let cache_dir = cache::hf_cache_dir()?;
let repo_dir = cache_layout::repo_dir(&cache_dir, repo_id);
let reader = HttpRangeReader::open(repo_id, revision, filename, token).await?;
let etag = reader.probe_etag().to_owned();
let cache_path = cache_layout::header_cache_path(&repo_dir, filename, &etag);
if let Some(entry) =
header_cache::HeaderCacheEntry::load(&cache_path, repo_id, rev, filename, &etag).await
{
let age = entry.cached_at.elapsed().unwrap_or_default();
return Ok((
entry.info,
inspect::InspectSource::CachedHeader { age },
None,
));
}
let (info, source, stats) =
dispatch_inspect_remote_from_reader(reader, filename, is_npz, is_gguf, is_pth).await?;
let entry = header_cache::HeaderCacheEntry::new(
repo_id.to_owned(),
rev.to_owned(),
filename.to_owned(),
etag,
info.clone(),
);
if let Err(e) = entry.save_atomic(&cache_path).await {
eprintln!("warning: failed to write header cache entry: {e}");
}
Ok((info, source, stats))
}
#[allow(
clippy::fn_params_excessive_bools,
clippy::too_many_arguments,
clippy::too_many_lines
)]
fn run_inspect_single(
repo_id: &str,
filename: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
cache_headers: bool,
no_metadata: bool,
json: bool,
filter: Option<&str>,
dtypes: bool,
group_by: Option<&str>,
limit: Option<usize>,
tree: bool,
check_gpu: Option<u32>,
context: Option<u32>,
) -> Result<(), FetchError> {
let ext_lc = Path::new(filename)
.extension()
.and_then(|e| e.to_str())
.map(str::to_ascii_lowercase);
let is_safetensors = ext_lc.as_deref() == Some("safetensors");
let is_gguf = ext_lc.as_deref() == Some("gguf");
let is_npz = ext_lc.as_deref() == Some("npz");
let is_pth = ext_lc.as_deref() == Some("pth");
if !is_safetensors && !is_gguf && !is_npz && !is_pth {
let extension = ext_lc.unwrap_or_else(|| "unknown".to_owned());
return Err(FetchError::UnsupportedInspectFormat {
filename: filename.to_owned(),
extension,
});
}
let (mut info, source, range_stats) = if cached {
let info = if is_gguf {
inspect::inspect_gguf_cached(repo_id, filename, revision)?
} else if is_npz {
inspect::inspect_npz_cached(repo_id, filename, revision)?
} else if is_pth {
inspect::inspect_pth_cached(repo_id, filename, revision)?
} else {
inspect::inspect_safetensors_cached(repo_id, filename, revision)?
};
(info, inspect::InspectSource::Cached, None)
} else {
let token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
rt.block_on(inspect_remote_with_cache(
repo_id,
filename,
revision,
token.as_deref(),
is_npz,
is_gguf,
is_pth,
cache_headers,
))?
};
let gpu_inputs = check_gpu.map(|idx| GpuCheckInputs {
device_index: idx,
weight_bytes: gpu_check::sum_tensor_bytes(&info.tensors),
dtype_label: gpu_check::dominant_dtype_label(&info.tensors),
total_params: info.total_params(),
kv: compute_kv_inputs(repo_id, revision, token, cached, context),
});
let total_tensor_count = info.tensors.len();
let total_params = info.total_params();
if let Some(pattern) = filter {
info.tensors
.retain(|t| matches_filter(t.name.as_str(), pattern));
}
let matched_count = info.tensors.len();
let matched_params = info.total_params();
let truncated_by_limit = limit.is_some_and(|n| matched_count > n);
if let Some(n) = limit {
info.tensors.truncate(n);
}
let gpu_result = gpu_inputs
.as_ref()
.map(|i| gpu_check::query_gpu(i.device_index));
let gpu_check_value = gpu_inputs.as_ref().zip(gpu_result.as_ref()).map(|(i, r)| {
gpu_check::gpu_check_json(
r,
i.weight_bytes,
i.dtype_label.as_str(),
i.total_params,
i.kv.as_ref(),
)
});
if tree && json {
return print_tree_json(
repo_id,
filename,
&info.tensors,
total_tensor_count,
total_params,
gpu_check_value,
);
}
if dtypes && json {
return print_dtype_summary_json(
&info.tensors,
total_tensor_count,
total_params,
gpu_check_value,
);
}
if let Some(pattern) = group_by
&& json
{
let matcher = compile_group_by_pattern(pattern)?;
let rollup = compute_group_by_rollup(&info.tensors, &matcher);
return print_group_by_summary_json(pattern, &rollup, gpu_check_value);
}
if json {
let wrapped = InspectJsonOutput {
header: &info,
truncated: truncated_by_limit.then_some(TruncationInfo {
shown: info.tensors.len(),
total: total_tensor_count,
}),
gpu_check: gpu_check_value,
};
let output = serde_json::to_string_pretty(&wrapped)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
return Ok(());
}
let source_label = match (source, range_stats) {
(inspect::InspectSource::Cached, _) => "cached".to_owned(),
(inspect::InspectSource::Remote, Some(stats)) => format!(
"remote ({} range requests, {} fetched)",
stats.requests,
format_size(stats.bytes_fetched)
),
(inspect::InspectSource::Remote, None) => "remote".to_owned(),
(inspect::InspectSource::CachedHeader { age }, _) => {
format!("cached header (age: {})", format_short_age(age))
}
_ => "unknown".to_owned(),
};
println!(" Repo: {repo_id}");
println!(" File: {filename}");
println!(" Source: {source_label}");
println!(
" {}",
format_header_line(
info.header_size,
info.file_size,
is_gguf || is_npz || is_pth
)
);
for line in format_quant_lines(info.quant_info.as_ref()) {
println!(" {line}");
}
if !no_metadata && let Some(ref meta) = info.metadata {
for line in format_metadata_lines(meta) {
println!(" {line}");
}
}
if tree {
print_tree_summary(&info.tensors, filter, total_tensor_count, total_params);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
if dtypes {
print_dtype_summary(&info.tensors, filter, total_tensor_count, total_params);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
if let Some(pattern) = group_by {
let matcher = compile_group_by_pattern(pattern)?;
let rollup = compute_group_by_rollup(&info.tensors, &matcher);
print_group_by_summary(pattern, &rollup);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
let nw = info
.tensors
.iter()
.map(|t| t.name.len())
.max()
.unwrap_or(6)
.max(6); let shape_strs: Vec<String> = info
.tensors
.iter()
.map(|t| format!("{:?}", t.shape))
.collect();
let sw = shape_strs.iter().map(String::len).max().unwrap_or(5).max(5); let row_width = nw + 2 + 8 + sw + 2 + 10 + 2 + 10;
println!();
println!(
" {:<nw$} {:<8} {:<sw$} {:>10} {:>10}",
"Tensor", "Dtype", "Shape", "Size", "Params",
);
for (t, shape_str) in info.tensors.iter().zip(shape_strs.iter()) {
let size_str = format_size(t.byte_len());
let params_str = inspect::format_params(t.num_elements());
println!(
" {:<nw$} {:<8} {:<sw$} {:>10} {:>10}",
t.name, t.dtype, shape_str, size_str, params_str,
);
}
println!(" {}", "\u{2500}".repeat(row_width));
let shown_count = info.tensors.len();
let shown_params = info.total_params();
let tensor_label = if shown_count == 1 {
"tensor"
} else {
"tensors"
};
match (filter.is_some(), truncated_by_limit) {
(false, false) => {
println!(
" {shown_count} {tensor_label}, {} params",
inspect::format_params(shown_params)
);
}
(true, false) => {
let filter_str = filter.unwrap_or_default();
println!(
" Showing {shown_count} of {total_tensor_count} tensors matching filter {filter_str:?}."
);
println!(
" Param counts: {} matching filter, {} total.",
inspect::format_params(shown_params),
inspect::format_params(total_params),
);
}
(false, true) => {
let limit_val = limit.unwrap_or(0);
println!(
" Showing {shown_count} of {total_tensor_count} tensors (limit: {limit_val})."
);
println!(
" Param counts: {} shown, {} total.",
inspect::format_params(shown_params),
inspect::format_params(total_params),
);
}
(true, true) => {
let filter_str = filter.unwrap_or_default();
let limit_val = limit.unwrap_or(0);
println!(
" Showing {shown_count} of {matched_count} tensors matching filter {filter_str:?} ({total_tensor_count} tensors total, limit: {limit_val})."
);
println!(
" Param counts: {} shown, {} matching filter, {} total.",
inspect::format_params(shown_params),
inspect::format_params(matched_params),
inspect::format_params(total_params),
);
}
}
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
Ok(())
}
struct GpuCheckInputs {
device_index: u32,
weight_bytes: u64,
dtype_label: String,
total_params: u64,
kv: Option<gpu_check::KvComputed>,
}
fn maybe_print_gpu_check(
inputs: Option<&GpuCheckInputs>,
probe: Option<&gpu_check::GpuCheckResult>,
) {
if let (Some(i), Some(result)) = (inputs, probe) {
gpu_check::print_gpu_check(
result,
i.weight_bytes,
i.dtype_label.as_str(),
i.total_params,
i.kv.as_ref(),
);
}
}
fn torch_dtype_label(torch_dtype: Option<&str>) -> &'static str {
match torch_dtype {
Some("float16") => "FP16",
Some("float32" | "float") => "FP32",
Some("float8_e4m3fn" | "float8_e5m2") => "FP8",
_ => "BF16",
}
}
fn fetch_model_config_for_inspect(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
) -> Option<inspect::ModelConfig> {
if cached {
return inspect::fetch_model_config_cached(repo_id, revision)
.ok()
.flatten();
}
let rt = tokio::runtime::Runtime::new().ok()?;
rt.block_on(inspect::fetch_model_config(repo_id, token, revision))
.ok()
.flatten()
}
fn compute_kv_inputs(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
context: Option<u32>,
) -> Option<gpu_check::KvComputed> {
let ctx = context?;
let Some(cfg) = fetch_model_config_for_inspect(repo_id, revision, token, cached) else {
return Some(gpu_check::KvComputed {
context: ctx,
elem_bytes: 0,
dtype_label: String::new(),
bytes: None,
path: gpu_check::KvCachePath::Unavailable {
reason: "no config.json in repo",
},
});
};
let elem_bytes = inspect::torch_dtype_bytes(cfg.torch_dtype.as_deref());
let dtype_label = torch_dtype_label(cfg.torch_dtype.as_deref()).to_owned();
let (bytes, path) = gpu_check::kv_cache_bytes(&cfg, u64::from(ctx), elem_bytes);
Some(gpu_check::KvComputed {
context: ctx,
elem_bytes,
dtype_label,
bytes,
path,
})
}
#[allow(clippy::too_many_arguments, clippy::fn_params_excessive_bools)]
fn run_inspect_repo(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
json: bool,
filter: Option<&str>,
dtypes: bool,
group_by: Option<&str>,
limit: Option<usize>,
tree: bool,
check_gpu: Option<u32>,
context: Option<u32>,
) -> Result<(), FetchError> {
let needs_aggregation =
dtypes || group_by.is_some() || tree || limit.is_some() || check_gpu.is_some();
let use_shard_fast_path = !needs_aggregation && !json;
let kv = compute_kv_inputs(repo_id, revision, token, cached, context);
if cached {
if use_shard_fast_path
&& let Some(index) = inspect::fetch_shard_index_cached(repo_id, revision)?
{
print_shard_index_summary(repo_id, &index, filter);
print_adapter_config_if_present(repo_id, revision, None, true, json);
return Ok(());
}
let results = inspect::inspect_repo_safetensors_cached(repo_id, revision)?;
if results.is_empty() {
println!("No cached .safetensors files found for {repo_id}.");
println!("Hint: use `hf-fm list-files {repo_id}` to see available file types");
return Ok(());
}
if needs_aggregation {
run_inspect_repo_aggregated(
repo_id,
&results,
json,
filter,
dtypes,
group_by,
limit,
tree,
check_gpu,
kv.as_ref(),
)?;
print_adapter_config_if_present(repo_id, revision, None, true, json);
return Ok(());
}
if json {
print_multi_file_json(&results, filter)?;
print_adapter_config_if_present(repo_id, revision, None, true, true);
return Ok(());
}
print_multi_file_summary(repo_id, "cached", &results, filter);
print_adapter_config_if_present(repo_id, revision, None, true, false);
return Ok(());
}
let token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
if use_shard_fast_path {
let shard_index = rt.block_on(inspect::fetch_shard_index(
repo_id,
token.as_deref(),
revision,
))?;
if let Some(index) = shard_index {
print_shard_index_summary(repo_id, &index, filter);
print_adapter_config_if_present(repo_id, revision, token.as_deref(), false, json);
return Ok(());
}
}
let results = rt.block_on(inspect::inspect_repo_safetensors(
repo_id,
token.as_deref(),
revision,
))?;
if results.is_empty() {
println!("No .safetensors files found in {repo_id}.");
println!("Hint: use `hf-fm list-files {repo_id}` to see available file types");
return Ok(());
}
let sources: Vec<inspect::InspectSource> = results.iter().map(|(_, _, s)| *s).collect();
let source_label = multi_file_source_label(&sources);
let mapped: Vec<(String, inspect::SafetensorsHeaderInfo)> = results
.into_iter()
.map(|(name, info, _source)| (name, info))
.collect();
if needs_aggregation {
run_inspect_repo_aggregated(
repo_id,
&mapped,
json,
filter,
dtypes,
group_by,
limit,
tree,
check_gpu,
kv.as_ref(),
)?;
print_adapter_config_if_present(repo_id, revision, token.as_deref(), false, json);
return Ok(());
}
if json {
print_multi_file_json(&mapped, filter)?;
print_adapter_config_if_present(repo_id, revision, token.as_deref(), false, true);
return Ok(());
}
print_multi_file_summary(repo_id, source_label.as_str(), &mapped, filter);
print_adapter_config_if_present(repo_id, revision, token.as_deref(), false, false);
Ok(())
}
#[allow(
clippy::fn_params_excessive_bools,
clippy::too_many_arguments,
clippy::too_many_lines
)]
fn run_inspect_repo_aggregated(
repo_id: &str,
results: &[(String, inspect::SafetensorsHeaderInfo)],
json: bool,
filter: Option<&str>,
dtypes: bool,
group_by: Option<&str>,
limit: Option<usize>,
tree: bool,
check_gpu: Option<u32>,
kv: Option<&gpu_check::KvComputed>,
) -> Result<(), FetchError> {
let mut flat: Vec<(&str, &inspect::TensorInfo)> = Vec::new();
for (file, info) in results {
for t in &info.tensors {
flat.push((file.as_str(), t));
}
}
let total_tensor_count = flat.len();
let total_params: u64 = flat
.iter()
.map(|(_, t)| t.num_elements())
.fold(0u64, u64::saturating_add);
let gpu_inputs = check_gpu.map(|idx| {
let owned_unfiltered: Vec<inspect::TensorInfo> =
flat.iter().map(|(_, t)| (*t).clone()).collect();
GpuCheckInputs {
device_index: idx,
weight_bytes: gpu_check::sum_tensor_bytes(&owned_unfiltered),
dtype_label: gpu_check::dominant_dtype_label(&owned_unfiltered),
total_params,
kv: kv.cloned(),
}
});
if let Some(pattern) = filter {
flat.retain(|(_, t)| matches_filter(t.name.as_str(), pattern));
}
let matched_count = flat.len();
let matched_params: u64 = flat
.iter()
.map(|(_, t)| t.num_elements())
.fold(0u64, u64::saturating_add);
let truncated_by_limit = limit.is_some_and(|n| matched_count > n);
if let Some(n) = limit {
flat.truncate(n);
}
let tensors_owned: Vec<inspect::TensorInfo> = flat.iter().map(|(_, t)| (*t).clone()).collect();
let gpu_result = gpu_inputs
.as_ref()
.map(|i| gpu_check::query_gpu(i.device_index));
let gpu_check_value = gpu_inputs.as_ref().zip(gpu_result.as_ref()).map(|(i, r)| {
gpu_check::gpu_check_json(
r,
i.weight_bytes,
i.dtype_label.as_str(),
i.total_params,
i.kv.as_ref(),
)
});
if tree && json {
return print_tree_json(
repo_id,
"<all shards>",
&tensors_owned,
total_tensor_count,
total_params,
gpu_check_value,
);
}
if dtypes && json {
return print_dtype_summary_json(
&tensors_owned,
total_tensor_count,
total_params,
gpu_check_value,
);
}
if let Some(pattern) = group_by
&& json
{
let matcher = compile_group_by_pattern(pattern)?;
let rollup = compute_group_by_rollup(&tensors_owned, &matcher);
return print_group_by_summary_json(pattern, &rollup, gpu_check_value);
}
if json {
return print_multi_file_json_with_gpu_check(results, filter, gpu_check_value);
}
if !tree && !dtypes && group_by.is_none() && limit.is_none() {
let n_shards = results.len();
let shard_label = if n_shards == 1 { "shard" } else { "shards" };
print_multi_file_summary(
repo_id,
&format!("aggregated across {n_shards} {shard_label}"),
results,
filter,
);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
println!(" Repo: {repo_id}");
let n_shards = results.len();
let shard_label = if n_shards == 1 { "shard" } else { "shards" };
println!(" Source: aggregated across {n_shards} {shard_label}");
if tree {
print_tree_summary(&tensors_owned, filter, total_tensor_count, total_params);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
if dtypes {
print_dtype_summary(&tensors_owned, filter, total_tensor_count, total_params);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
if let Some(pattern) = group_by {
let matcher = compile_group_by_pattern(pattern)?;
let rollup = compute_group_by_rollup(&tensors_owned, &matcher);
print_group_by_summary(pattern, &rollup);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
return Ok(());
}
print_multi_shard_table(
&flat,
filter,
limit,
truncated_by_limit,
total_tensor_count,
matched_count,
total_params,
matched_params,
);
maybe_print_gpu_check(gpu_inputs.as_ref(), gpu_result.as_ref());
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn print_multi_shard_table(
flat: &[(&str, &inspect::TensorInfo)],
filter: Option<&str>,
limit: Option<usize>,
truncated_by_limit: bool,
total_tensor_count: usize,
matched_count: usize,
total_params: u64,
matched_params: u64,
) {
let nw = flat
.iter()
.map(|(_, t)| t.name.len())
.max()
.unwrap_or(6)
.max(6); let shape_strs: Vec<String> = flat.iter().map(|(_, t)| format!("{:?}", t.shape)).collect();
let sw = shape_strs.iter().map(String::len).max().unwrap_or(5).max(5); let fw = flat
.iter()
.map(|(file, _)| file.len())
.max()
.unwrap_or(5)
.max(5); let row_width = nw + 2 + 8 + 2 + sw + 2 + 10 + 2 + 10 + 2 + fw;
println!();
println!(
" {:<nw$} {:<8} {:<sw$} {:>10} {:>10} {:<fw$}",
"Tensor", "Dtype", "Shape", "Size", "Params", "Shard",
);
for ((file, t), shape_str) in flat.iter().zip(shape_strs.iter()) {
let size_str = format_size(t.byte_len());
let params_str = inspect::format_params(t.num_elements());
println!(
" {:<nw$} {:<8} {:<sw$} {:>10} {:>10} {:<fw$}",
t.name, t.dtype, shape_str, size_str, params_str, file,
);
}
println!(" {}", "\u{2500}".repeat(row_width));
let shown_count = flat.len();
let shown_params: u64 = flat
.iter()
.map(|(_, t)| t.num_elements())
.fold(0u64, u64::saturating_add);
let tensor_label = if shown_count == 1 {
"tensor"
} else {
"tensors"
};
match (filter.is_some(), truncated_by_limit) {
(false, false) => {
println!(
" {shown_count} {tensor_label}, {} params",
inspect::format_params(shown_params)
);
}
(true, false) => {
println!(
" {shown_count}/{total_tensor_count} {tensor_label}, {}/{} params (filter: {:?})",
inspect::format_params(shown_params),
inspect::format_params(total_params),
filter.unwrap_or_default(),
);
}
(false, true) => {
println!(
" {shown_count}/{total_tensor_count} {tensor_label} shown, {}/{} params (limit: {})",
inspect::format_params(shown_params),
inspect::format_params(total_params),
limit.unwrap_or(0),
);
}
(true, true) => {
println!(
" {shown_count}/{matched_count}/{total_tensor_count} {tensor_label} shown, {}/{}/{} params (filter: {:?}, limit: {})",
inspect::format_params(shown_params),
inspect::format_params(matched_params),
inspect::format_params(total_params),
filter.unwrap_or_default(),
limit.unwrap_or(0),
);
}
}
}
fn print_adapter_config_if_present(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
cached: bool,
json: bool,
) {
let result = if cached {
inspect::fetch_adapter_config_cached(repo_id, revision)
} else {
let Ok(rt) = tokio::runtime::Runtime::new() else {
return;
};
rt.block_on(inspect::fetch_adapter_config(repo_id, token, revision))
};
let Ok(Some(config)) = result else { return };
if json {
if let Ok(output) = serde_json::to_string_pretty(&config) {
println!("{output}");
}
return;
}
println!();
println!(" Adapter config:");
if let Some(ref peft_type) = config.peft_type {
println!(" PEFT type: {peft_type}");
}
if let Some(ref base) = config.base_model_name_or_path {
println!(" Base model: {base}");
}
if let Some(r) = config.r {
println!(" Rank (r): {r}");
}
if let Some(alpha) = config.lora_alpha {
println!(" LoRA alpha: {alpha}");
}
if let Some(ref task) = config.task_type {
println!(" Task type: {task}");
}
if !config.target_modules.is_empty() {
println!(" Target modules: {}", config.target_modules.join(", "));
}
}
#[derive(serde::Serialize)]
struct DtypeGroup<'a> {
dtype: &'a str,
tensors: usize,
params: u64,
bytes: u64,
}
#[derive(serde::Serialize)]
struct DtypeSummaryJson<'a> {
dtypes: Vec<DtypeGroup<'a>>,
total_tensors: usize,
total_params: u64,
#[serde(skip_serializing_if = "Option::is_none")]
gpu_check: Option<serde_json::Value>,
}
fn compute_dtype_groups<'a, I>(tensors: I) -> Vec<(&'a str, usize, u64, u64)>
where
I: IntoIterator<Item = &'a inspect::TensorInfo>,
{
let mut groups: HashMap<&str, (usize, u64, u64)> = HashMap::new();
for t in tensors {
let entry = groups
.entry(t.dtype.as_str()) .or_insert((0, 0, 0));
entry.0 += 1;
entry.1 = entry.1.saturating_add(t.num_elements());
entry.2 = entry.2.saturating_add(t.byte_len());
}
let mut rows: Vec<(&str, usize, u64, u64)> = groups
.into_iter()
.map(|(dtype, (count, params, bytes))| (dtype, count, params, bytes))
.collect();
rows.sort_by_key(|r| std::cmp::Reverse(r.1));
rows
}
fn print_dtype_summary_json(
tensors: &[inspect::TensorInfo],
total_tensor_count: usize,
total_params: u64,
gpu_check: Option<serde_json::Value>,
) -> Result<(), FetchError> {
let rows = compute_dtype_groups(tensors);
let output = DtypeSummaryJson {
dtypes: rows
.into_iter()
.map(|(dtype, tensors, params, bytes)| DtypeGroup {
dtype,
tensors,
params,
bytes,
})
.collect(),
total_tensors: total_tensor_count,
total_params,
gpu_check,
};
let serialized = serde_json::to_string_pretty(&output)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{serialized}");
Ok(())
}
fn print_dtype_summary(
tensors: &[inspect::TensorInfo],
filter: Option<&str>,
total_tensor_count: usize,
total_params: u64,
) {
let rows = compute_dtype_groups(tensors);
let dw = rows
.iter()
.map(|(d, _, _, _)| d.len())
.max()
.unwrap_or(5)
.max(5); let row_width = dw + 2 + 8 + 2 + 12 + 2 + 10;
println!();
println!(
" {:<dw$} {:>8} {:>12} {:>10}",
"Dtype", "Tensors", "Params", "Size",
);
for (dtype, count, params, bytes) in &rows {
println!(
" {:<dw$} {:>8} {:>12} {:>10}",
dtype,
count,
inspect::format_params(*params),
format_size(*bytes),
);
}
println!(" {}", "\u{2500}".repeat(row_width));
let filtered_count: usize = rows.iter().map(|(_, count, _, _)| count).sum();
let filtered_params: u64 = rows.iter().map(|(_, _, params, _)| params).sum();
let tensor_label = if filtered_count == 1 {
"tensor"
} else {
"tensors"
};
if filter.is_some() {
println!(
" {filtered_count}/{total_tensor_count} {tensor_label}, {}/{} params",
inspect::format_params(filtered_params),
inspect::format_params(total_params),
);
} else {
println!(
" {filtered_count} {tensor_label}, {} params",
inspect::format_params(filtered_params),
);
}
}
fn compile_group_by_pattern(pattern: &str) -> Result<globset::GlobMatcher, FetchError> {
globset::Glob::new(pattern)
.map(|g| g.compile_matcher())
.map_err(|e| FetchError::InvalidPattern {
pattern: pattern.to_owned(),
reason: e.to_string(),
})
}
fn distinct_layer_indices(names: &[&str]) -> Option<Vec<usize>> {
let segment_count = names.first()?.split('.').count();
let split: Vec<Vec<&str>> = names.iter().map(|n| n.split('.').collect()).collect();
if split.iter().any(|s| s.len() != segment_count) {
return None;
}
let mut varying: Option<Vec<usize>> = None;
for pos in 0..segment_count {
let parsed: Option<Vec<usize>> = split
.iter()
.map(|s| s.get(pos).and_then(|seg| seg.parse::<usize>().ok()))
.collect();
let Some(values) = parsed else {
continue; };
let unique: std::collections::BTreeSet<usize> = values.iter().copied().collect();
if unique.len() <= 1 {
continue; }
if varying.is_some() {
return None; }
varying = Some(unique.into_iter().collect());
}
varying
}
struct GroupByRollup {
matched_tensors: usize,
matched_params: u64,
matched_bytes: u64,
other_tensors: usize,
other_params: u64,
other_bytes: u64,
total_tensors: usize,
total_params: u64,
total_bytes: u64,
layer_count: Option<usize>,
per_layer_bytes: Option<u64>,
}
fn compute_group_by_rollup<'a, I>(tensors: I, matcher: &globset::GlobMatcher) -> GroupByRollup
where
I: IntoIterator<Item = &'a inspect::TensorInfo>,
{
let mut matched_tensors = 0usize;
let mut matched_params = 0u64;
let mut matched_bytes = 0u64;
let mut other_tensors = 0usize;
let mut other_params = 0u64;
let mut other_bytes = 0u64;
let mut matched_names: Vec<&str> = Vec::new();
for t in tensors {
if matcher.is_match(&t.name) {
matched_tensors += 1;
matched_params = matched_params.saturating_add(t.num_elements());
matched_bytes = matched_bytes.saturating_add(t.byte_len());
matched_names.push(t.name.as_str()); } else {
other_tensors += 1;
other_params = other_params.saturating_add(t.num_elements());
other_bytes = other_bytes.saturating_add(t.byte_len());
}
}
let layer_count = distinct_layer_indices(&matched_names).map(|v| v.len());
let per_layer_bytes = layer_count.and_then(|n| {
if n == 0 {
None
} else {
#[allow(clippy::as_conversions)]
let n = n as u64;
Some(matched_bytes / n)
}
});
GroupByRollup {
matched_tensors,
matched_params,
matched_bytes,
other_tensors,
other_params,
other_bytes,
total_tensors: matched_tensors + other_tensors,
total_params: matched_params.saturating_add(other_params),
total_bytes: matched_bytes.saturating_add(other_bytes),
layer_count,
per_layer_bytes,
}
}
fn print_group_by_summary(pattern: &str, rollup: &GroupByRollup) {
let matched_label = format!("MATCHED ({pattern})");
let label_width = matched_label.len().max("OTHER".len()).max("TOTAL".len());
let row_width = label_width + 2 + 8 + 2 + 12 + 2 + 8;
let pct = |bytes: u64| -> f64 {
if rollup.total_bytes == 0 {
0.0
} else {
#[allow(clippy::cast_precision_loss, clippy::as_conversions)]
let ratio = bytes as f64 / rollup.total_bytes as f64;
ratio * 100.0
}
};
println!();
println!(
" {:<label_width$} {:>8} {:>12} {:>8}",
"Group", "Tensors", "Size", "Percent"
);
println!(
" {:<label_width$} {:>8} {:>12} {:>7.1}%",
matched_label,
rollup.matched_tensors,
format_size(rollup.matched_bytes),
pct(rollup.matched_bytes),
);
println!(
" {:<label_width$} {:>8} {:>12} {:>7.1}%",
"OTHER",
rollup.other_tensors,
format_size(rollup.other_bytes),
pct(rollup.other_bytes),
);
println!(" {}", "\u{2500}".repeat(row_width));
println!(
" {:<label_width$} {:>8} {:>12} {:>7.1}%",
"TOTAL",
rollup.total_tensors,
format_size(rollup.total_bytes),
pct(rollup.total_bytes),
);
if let (Some(layer_count), Some(per_layer)) = (rollup.layer_count, rollup.per_layer_bytes) {
println!();
println!(
" per-MoE-layer expert cost: {} ({layer_count} layers)",
format_size(per_layer)
);
}
}
#[derive(serde::Serialize)]
struct GroupByBucketJson {
tensors: usize,
params: u64,
bytes: u64,
}
#[derive(serde::Serialize)]
struct GroupByJson<'a> {
pattern: &'a str,
matched: GroupByBucketJson,
other: GroupByBucketJson,
total_tensors: usize,
total_params: u64,
total_bytes: u64,
#[serde(skip_serializing_if = "Option::is_none")]
layer_count: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
per_layer_bytes: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
gpu_check: Option<serde_json::Value>,
}
fn print_group_by_summary_json(
pattern: &str,
rollup: &GroupByRollup,
gpu_check: Option<serde_json::Value>,
) -> Result<(), FetchError> {
let output = GroupByJson {
pattern,
matched: GroupByBucketJson {
tensors: rollup.matched_tensors,
params: rollup.matched_params,
bytes: rollup.matched_bytes,
},
other: GroupByBucketJson {
tensors: rollup.other_tensors,
params: rollup.other_params,
bytes: rollup.other_bytes,
},
total_tensors: rollup.total_tensors,
total_params: rollup.total_params,
total_bytes: rollup.total_bytes,
layer_count: rollup.layer_count,
per_layer_bytes: rollup.per_layer_bytes,
gpu_check,
};
let serialized = serde_json::to_string_pretty(&output)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{serialized}");
Ok(())
}
fn print_shard_index_summary(repo_id: &str, index: &inspect::ShardedIndex, filter: Option<&str>) {
println!(" Repo: {repo_id}");
println!(" Source: shard index (model.safetensors.index.json)");
println!();
let total_tensors = index.weight_map.len();
let mut by_shard: HashMap<&str, usize> = HashMap::new();
let mut names_by_shard: HashMap<&str, Vec<&str>> = HashMap::new();
let mut filtered_total: usize = 0;
for (tensor_name, shard_name) in &index.weight_map {
if let Some(pattern) = filter {
if !matches_filter(tensor_name, pattern) {
continue;
}
names_by_shard
.entry(shard_name.as_str())
.or_default()
.push(tensor_name.as_str());
}
*by_shard.entry(shard_name.as_str()).or_default() += 1;
filtered_total += 1;
}
let fw = index
.shards
.iter()
.map(String::len)
.max()
.unwrap_or(4)
.max(4); let row_width = fw + 2 + 8;
println!(" {:<fw$} {:>8}", "File", "Tensors");
for shard in &index.shards {
let count = by_shard.get(shard.as_str()).copied().unwrap_or(0);
if filter.is_some() && count == 0 {
continue;
}
println!(" {shard:<fw$} {count:>8}");
if filter.is_some()
&& let Some(names) = names_by_shard.get(shard.as_str())
{
let mut sorted = names.clone();
sorted.sort_unstable();
for tname in sorted {
println!(" {tname}");
}
}
}
println!(" {}", "\u{2500}".repeat(row_width));
let displayed_shards = if filter.is_some() {
by_shard.len()
} else {
index.shards.len()
};
let shard_label = if displayed_shards == 1 {
"shard"
} else {
"shards"
};
let tensor_label = if filtered_total == 1 {
"tensor"
} else {
"tensors"
};
if filter.is_some() {
println!(
" {displayed_shards} {shard_label}, {filtered_total}/{total_tensors} {tensor_label} (filter: {:?})",
filter.unwrap_or_default(),
);
println!(
" Hint: names shown above \u{2014} for shapes/dtypes/sizes run \
`hf-fm inspect {repo_id} <filename> --tree` (or `--dtypes`), or add \
`--limit N` to cap a broad match."
);
} else {
println!(" {displayed_shards} {shard_label}, {filtered_total} {tensor_label}");
println!(
" Hint: this rollup hides tensor names \u{2014} run \
`hf-fm inspect {repo_id} <filename> --tree` (or `--dtypes`) for per-tensor detail."
);
}
}
fn print_multi_file_json(
results: &[(String, inspect::SafetensorsHeaderInfo)],
filter: Option<&str>,
) -> Result<(), FetchError> {
if let Some(pattern) = filter {
let filtered: Vec<(String, inspect::SafetensorsHeaderInfo)> = results
.iter()
.filter_map(|(name, info)| {
let matching: Vec<inspect::TensorInfo> = info
.tensors
.iter()
.filter(|t| matches_filter(t.name.as_str(), pattern)) .cloned()
.collect();
if matching.is_empty() {
return None;
}
Some((
name.clone(), inspect::SafetensorsHeaderInfo::new(
matching,
info.metadata.clone(),
info.header_size,
info.file_size,
info.quant_info.clone(),
),
))
})
.collect();
let output = serde_json::to_string_pretty(&filtered)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
} else {
let output = serde_json::to_string_pretty(results)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
}
Ok(())
}
fn print_multi_file_json_with_gpu_check(
results: &[(String, inspect::SafetensorsHeaderInfo)],
filter: Option<&str>,
gpu_check: Option<serde_json::Value>,
) -> Result<(), FetchError> {
let files_payload: Vec<(String, inspect::SafetensorsHeaderInfo)> = if let Some(pattern) = filter
{
results
.iter()
.filter_map(|(name, info)| {
let matching: Vec<inspect::TensorInfo> = info
.tensors
.iter()
.filter(|t| matches_filter(t.name.as_str(), pattern))
.cloned()
.collect();
if matching.is_empty() {
return None;
}
Some((
name.clone(),
inspect::SafetensorsHeaderInfo::new(
matching,
info.metadata.clone(),
info.header_size,
info.file_size,
info.quant_info.clone(),
),
))
})
.collect()
} else {
results.to_vec()
};
let mut top = serde_json::Map::new();
top.insert(
"files".to_owned(),
serde_json::to_value(&files_payload)
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?,
);
if let Some(gc) = gpu_check {
top.insert("gpu_check".to_owned(), gc);
}
let output = serde_json::to_string_pretty(&serde_json::Value::Object(top))
.map_err(|e| FetchError::Http(format!("failed to serialize JSON: {e}")))?;
println!("{output}");
Ok(())
}
fn multi_file_source_label(sources: &[inspect::InspectSource]) -> String {
let cached = sources
.iter()
.filter(|s| matches!(s, inspect::InspectSource::Cached))
.count();
let remote = sources
.iter()
.filter(|s| matches!(s, inspect::InspectSource::Remote))
.count();
match (cached, remote) {
(0, 0) => "unknown".to_owned(),
(_, 0) => "cached".to_owned(),
(0, _) => "remote".to_owned(),
(c, r) => format!("mixed ({c} cached, {r} remote)"),
}
}
fn print_multi_file_summary(
repo_id: &str,
source: &str,
results: &[(String, inspect::SafetensorsHeaderInfo)],
filter: Option<&str>,
) {
println!(" Repo: {repo_id}");
println!(" Source: {source}");
println!();
let fw = results
.iter()
.map(|(name, _)| name.len())
.max()
.unwrap_or(4)
.max(4); let row_width = fw + 2 + 8 + 1 + 12;
println!(" {:<fw$} {:>8} {:>12}", "File", "Tensors", "Params");
let mut total_tensors_unfiltered: usize = 0;
let mut total_params_unfiltered: u64 = 0;
let mut total_tensors_filtered: usize = 0;
let mut total_params_filtered: u64 = 0;
let mut files_with_matches: usize = 0;
for (name, info) in results {
total_tensors_unfiltered = total_tensors_unfiltered.saturating_add(info.tensors.len());
total_params_unfiltered = total_params_unfiltered.saturating_add(info.total_params());
let (tensor_count, params, matched_names) = if let Some(pattern) = filter {
let matching: Vec<&inspect::TensorInfo> = info
.tensors
.iter()
.filter(|t| matches_filter(t.name.as_str(), pattern))
.collect();
let p: u64 = matching.iter().map(|t| t.num_elements()).sum();
let mut names: Vec<&str> = matching.iter().map(|t| t.name.as_str()).collect();
names.sort_unstable();
(matching.len(), p, names)
} else {
(info.tensors.len(), info.total_params(), Vec::new())
};
if filter.is_some() && tensor_count == 0 {
continue;
}
files_with_matches += 1;
total_tensors_filtered = total_tensors_filtered.saturating_add(tensor_count);
total_params_filtered = total_params_filtered.saturating_add(params);
println!(
" {name:<fw$} {tensor_count:>8} {:>12}",
inspect::format_params(params)
);
for tname in &matched_names {
println!(" {tname}");
}
}
println!(" {}", "\u{2500}".repeat(row_width));
let file_label = if files_with_matches == 1 {
"file"
} else {
"files"
};
let tensor_label = if total_tensors_filtered == 1 {
"tensor"
} else {
"tensors"
};
if filter.is_some() {
println!(
" {} {file_label}, {total_tensors_filtered}/{total_tensors_unfiltered} {tensor_label}, {}/{} params (filter: {:?})",
files_with_matches,
inspect::format_params(total_params_filtered),
inspect::format_params(total_params_unfiltered),
filter.unwrap_or_default(),
);
} else {
println!(
" {} {file_label}, {total_tensors_filtered} {tensor_label}, {} params",
files_with_matches,
inspect::format_params(total_params_filtered)
);
}
if filter.is_some() {
println!(
" Hint: names shown above \u{2014} for shapes/dtypes/sizes run \
`hf-fm inspect {repo_id} <filename> --tree` (or `--dtypes`), or add \
`--limit N` to cap a broad match."
);
} else if let [(name, _)] = results {
println!(
" Hint: this rollup hides tensor names \u{2014} run \
`hf-fm inspect {repo_id} {name} --tree` (or `--dtypes`) for per-tensor detail."
);
}
}
#[allow(clippy::too_many_lines)]
fn run_status(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
preset: Option<&Preset>,
json: bool,
) -> Result<(), FetchError> {
let token = token
.map(String::from)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let cache_root = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_root, repo_id);
let sidecar = cache::read_snapshot(&repo_dir)?;
let effective_preset_name: Option<String> = match preset {
Some(p) => Some(preset_name(p).to_owned()), None => sidecar.and_then(|s| s.preset),
};
let preset_glob_list: Option<&'static [&'static str]> = effective_preset_name
.as_deref() .and_then(hf_fetch_model::config::preset_globs);
let status = rt.block_on(cache::repo_status(
repo_id,
token.as_deref(),
revision,
preset_glob_list,
))?;
if json {
return print_status_json(&status, revision.unwrap_or("main"));
}
let rev_display = revision.unwrap_or("main");
match &status.commit_hash {
Some(hash) => println!("{repo_id} ({rev_display} @ {hash})"),
None => println!("{repo_id} ({rev_display}, not yet cached)"),
}
println!("Cache: {}\n", status.cache_path.display());
if status.files.is_empty() {
println!(" (no files found in remote repository)");
return Ok(());
}
let fw = status
.files
.iter()
.map(|(name, _)| name.len())
.max()
.unwrap_or(4)
.max(4); for (filename, file_status) in &status.files {
match file_status {
cache::FileStatus::Complete { local_size } => {
println!(
" {:<fw$} {:>10} complete",
filename,
format_size(*local_size)
);
}
cache::FileStatus::Partial {
local_size,
expected_size,
} => {
println!(
" {:<fw$} {:>10} / {:<10} PARTIAL",
filename,
format_size(*local_size),
format_size(*expected_size)
);
}
cache::FileStatus::Missing { expected_size } => {
if *expected_size > 0 {
println!(
" {:<fw$} {:>10} MISSING",
filename,
format_size(*expected_size)
);
} else {
println!(" {filename:<fw$} {:>10} MISSING", "\u{2014}");
}
}
cache::FileStatus::Excluded { expected_size } => {
if *expected_size > 0 {
println!(
" {:<fw$} {:>10} excluded",
filename,
format_size(*expected_size)
);
} else {
println!(" {filename:<fw$} {:>10} excluded", "\u{2014}");
}
}
_ => {
println!(" {filename:<fw$} UNKNOWN");
}
}
}
let total = status.files.len();
let complete = status.complete_count();
let partial = status.partial_count();
let missing = status.missing_count();
let excluded = status.excluded_count();
println!();
if excluded > 0 {
println!(
"{complete}/{total} complete, {partial} partial, {missing} missing, {excluded} excluded"
);
} else {
println!("{complete}/{total} complete, {partial} partial, {missing} missing");
}
Ok(())
}
#[derive(serde::Serialize)]
struct StatusRepoSummaryJson {
repo_id: String,
file_count: usize,
size: u64,
has_partial: bool,
}
#[derive(serde::Serialize)]
struct StatusAllJson {
cache_dir: String,
repos: Vec<StatusRepoSummaryJson>,
model_count: usize,
}
#[derive(serde::Serialize)]
struct StatusFileJson {
filename: String,
state: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
local_size: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
expected_size: Option<u64>,
}
#[derive(serde::Serialize)]
struct StatusSummaryJson {
total: usize,
complete: usize,
partial: usize,
missing: usize,
excluded: usize,
}
#[derive(serde::Serialize)]
struct StatusRepoJson {
repo_id: String,
revision: String,
#[serde(skip_serializing_if = "Option::is_none")]
commit_hash: Option<String>,
cache_path: String,
files: Vec<StatusFileJson>,
summary: StatusSummaryJson,
}
fn print_status_all_json(
summaries: &[cache::CachedModelSummary],
cache_dir: &std::path::Path,
) -> Result<(), FetchError> {
let repos: Vec<StatusRepoSummaryJson> = summaries
.iter()
.map(|s| StatusRepoSummaryJson {
repo_id: s.repo_id.clone(),
file_count: s.file_count,
size: s.total_size,
has_partial: s.has_partial,
})
.collect();
let result = StatusAllJson {
cache_dir: cache_dir.display().to_string(),
model_count: repos.len(),
repos,
};
emit_json(&result)
}
fn file_status_json_fields(status: &cache::FileStatus) -> (&'static str, Option<u64>, Option<u64>) {
match status {
cache::FileStatus::Complete { local_size } => ("complete", Some(*local_size), None),
cache::FileStatus::Partial {
local_size,
expected_size,
} => ("partial", Some(*local_size), Some(*expected_size)),
cache::FileStatus::Missing { expected_size } => ("missing", None, Some(*expected_size)),
cache::FileStatus::Excluded { expected_size } => ("excluded", None, Some(*expected_size)),
_ => ("unknown", None, None),
}
}
fn print_status_json(status: &cache::RepoStatus, revision: &str) -> Result<(), FetchError> {
let files: Vec<StatusFileJson> = status
.files
.iter()
.map(|(filename, fs)| {
let (state, local_size, expected_size) = file_status_json_fields(fs);
StatusFileJson {
filename: filename.clone(),
state,
local_size,
expected_size,
}
})
.collect();
let summary = StatusSummaryJson {
total: status.files.len(),
complete: status.complete_count(),
partial: status.partial_count(),
missing: status.missing_count(),
excluded: status.excluded_count(),
};
let result = StatusRepoJson {
repo_id: status.repo_id.clone(),
revision: revision.to_owned(),
commit_hash: status.commit_hash.clone(),
cache_path: status.cache_path.display().to_string(),
files,
summary,
};
emit_json(&result)
}
#[derive(Clone, Copy)]
enum FileCacheState {
Complete,
Partial,
Missing,
}
impl FileCacheState {
const fn glyph(self) -> &'static str {
match self {
Self::Complete => "\u{2713}",
Self::Partial => "partial",
Self::Missing => "\u{2717}",
}
}
const fn word(self) -> &'static str {
match self {
Self::Complete => "complete",
Self::Partial => "partial",
Self::Missing => "missing",
}
}
}
#[allow(
clippy::too_many_arguments,
clippy::fn_params_excessive_bools,
clippy::too_many_lines
)]
fn run_list_files(
repo_id: &str,
revision: Option<&str>,
token: Option<&str>,
filter_patterns: &[String],
exclude_patterns: &[String],
preset: Option<&Preset>,
no_checksum: bool,
show_cached: bool,
json: bool,
) -> Result<(), FetchError> {
if !repo_id.contains('/') {
return Err(FetchError::InvalidArgument(format!(
"invalid REPO_ID \"{repo_id}\": expected \"org/model\" format \
(e.g., \"google/gemma-2-2b-it\")"
)));
}
let mut include_patterns: Vec<String> = match preset {
Some(&Preset::Safetensors) => vec![
"*.safetensors".to_owned(),
"*.json".to_owned(),
"*.txt".to_owned(),
],
Some(&Preset::Gguf) => vec!["*.gguf".to_owned(), "*.json".to_owned(), "*.txt".to_owned()],
Some(&Preset::Npz) => vec![
"*.npz".to_owned(),
"*.npy".to_owned(),
"config.yaml".to_owned(),
"*.json".to_owned(),
"*.txt".to_owned(),
],
Some(&Preset::Pth) => vec![
"pytorch_model*.bin".to_owned(),
"*.json".to_owned(),
"*.txt".to_owned(),
],
Some(&Preset::ConfigOnly) => {
vec!["*.json".to_owned(), "*.txt".to_owned(), "*.md".to_owned()]
}
None => Vec::new(),
};
for p in filter_patterns {
include_patterns.push(p.clone());
}
let include = compile_glob_patterns(&include_patterns)?;
let exclude = compile_glob_patterns(exclude_patterns)?;
let resolved_token = token
.map(ToOwned::to_owned)
.or_else(|| std::env::var("HF_TOKEN").ok());
let rt = tokio::runtime::Runtime::new().map_err(|e| FetchError::Io {
path: PathBuf::from("<runtime>"),
source: e,
})?;
let client = hf_fetch_model::build_client(resolved_token.as_deref())?;
let files = rt.block_on(repo::list_repo_files_with_metadata(
repo_id,
resolved_token.as_deref(),
revision,
&client,
))?;
let filtered: Vec<_> = files
.into_iter()
.filter(|f| {
file_matches(f.filename.as_str(), include.as_ref(), exclude.as_ref())
})
.collect();
let cache_marks: Vec<FileCacheState> = if show_cached {
let cache_dir = cache::hf_cache_dir()?;
let repo_dir = hf_fetch_model::cache_layout::repo_dir(&cache_dir, repo_id);
let revision_str = revision.unwrap_or("main");
let commit_hash = cache::read_ref(&repo_dir, revision_str);
let snapshot_dir =
commit_hash.map(|h| hf_fetch_model::cache_layout::snapshot_dir(&repo_dir, &h));
filtered
.iter()
.map(|f| {
let local_path = snapshot_dir
.as_ref()
.map(|dir| dir.join(f.filename.as_str()));
match local_path {
Some(ref path) if path.exists() => {
let local_size = std::fs::metadata(path).map_or(0, |m| m.len());
let expected = f.size.unwrap_or(0);
if expected > 0 && local_size < expected {
FileCacheState::Partial
} else {
FileCacheState::Complete
}
}
_ => FileCacheState::Missing,
}
})
.collect()
} else {
Vec::new()
};
if json {
return print_list_files_json(repo_id, &filtered, &cache_marks, show_cached);
}
let fw = filtered
.iter()
.map(|f| f.filename.len())
.max()
.unwrap_or(4)
.max(4);
if no_checksum {
if show_cached {
println!(" {:<fw$} {:>10} Cached", "File", "Size");
println!(" {:<fw$} {:>10} {:-<6}", "", "", "");
} else {
println!(" {:<fw$} {:>10}", "File", "Size");
println!(" {:<fw$} {:>10}", "", "");
}
} else if show_cached {
println!(" {:<fw$} {:>10} {:<12} Cached", "File", "Size", "SHA256");
println!(" {:<fw$} {:>10} {:<12} {:-<6}", "", "", "", "");
} else {
println!(" {:<fw$} {:>10} {:<12}", "File", "Size", "SHA256");
println!(" {:<fw$} {:>10} {:<12}", "", "", "");
}
let mut total_bytes: u64 = 0;
let mut cached_count: usize = 0;
let mut any_no_sha = false;
for (i, f) in filtered.iter().enumerate() {
let size = f.size.unwrap_or(0);
total_bytes = total_bytes.saturating_add(size);
let size_str = format_size(size);
let sha_str = if no_checksum {
String::new()
} else if let Some(hash) = f.sha256.as_deref().and_then(|s| s.get(..12)) {
hash.to_owned() } else {
any_no_sha = true;
"\u{2014}".to_owned()
};
if show_cached {
let state = cache_marks
.get(i)
.copied()
.unwrap_or(FileCacheState::Missing);
if matches!(state, FileCacheState::Complete) {
cached_count += 1;
}
let mark = state.glyph();
if no_checksum {
println!(" {:<fw$} {:>10} {mark}", f.filename, size_str);
} else {
println!(
" {:<fw$} {:>10} {:<12} {mark}",
f.filename, size_str, sha_str
);
}
} else if no_checksum {
println!(" {:<fw$} {:>10}", f.filename, size_str);
} else {
println!(" {:<fw$} {:>10} {sha_str}", f.filename, size_str);
}
}
let count = filtered.len();
let file_label = pluralize(count, "file", "files");
let row_width = fw + 2 + 10 + 2 + 12;
println!(" {:\u{2500}<row_width$}", "");
let sized: Vec<(&str, Option<u64>)> = filtered
.iter()
.map(|f| (f.filename.as_str(), f.size))
.collect();
if let Some((min, max)) = discover::gguf_size_range(sized) {
let cached_suffix = if show_cached {
format!(" ({cached_count} cached)")
} else {
String::new()
};
println!(
" {count} {file_label}, {} to {} (mutually exclusive quants){cached_suffix}",
format_size(min),
format_size(max),
);
} else if show_cached {
println!(
" {count} {file_label}, {} total ({cached_count} cached)",
format_size(total_bytes)
);
} else {
println!(" {count} {file_label}, {} total", format_size(total_bytes));
}
if any_no_sha && !no_checksum {
println!(" \u{2014} = not an LFS file (no SHA256 tracked by the Hub)");
}
Ok(())
}
#[derive(serde::Serialize)]
struct ListFileEntry {
filename: String,
size: u64,
sha256: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
cached: Option<String>,
}
#[derive(serde::Serialize)]
struct ListFilesResult {
repo_id: String,
files: Vec<ListFileEntry>,
total_bytes: u64,
file_count: usize,
#[serde(skip_serializing_if = "Option::is_none")]
cached_count: Option<usize>,
quant_alternatives: bool,
#[serde(skip_serializing_if = "Option::is_none")]
size_min: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
size_max: Option<u64>,
}
fn print_list_files_json(
repo_id: &str,
files: &[repo::RepoFile],
cache_marks: &[FileCacheState],
show_cached: bool,
) -> Result<(), FetchError> {
let mut entries: Vec<ListFileEntry> = Vec::with_capacity(files.len());
let mut total_bytes: u64 = 0;
let mut cached_count: usize = 0;
for (i, f) in files.iter().enumerate() {
let size = f.size.unwrap_or(0);
total_bytes = total_bytes.saturating_add(size);
let cached = if show_cached {
let state = cache_marks
.get(i)
.copied()
.unwrap_or(FileCacheState::Missing);
if matches!(state, FileCacheState::Complete) {
cached_count += 1;
}
Some(state.word().to_owned())
} else {
None
};
entries.push(ListFileEntry {
filename: f.filename.clone(),
size,
sha256: f.sha256.clone(),
cached,
});
}
let sized: Vec<(&str, Option<u64>)> = files
.iter()
.map(|f| (f.filename.as_str(), f.size))
.collect();
let range = discover::gguf_size_range(sized);
let quant_alternatives = range.is_some();
let size_min = range.map(|(min, _)| min);
let size_max = range.map(|(_, max)| max);
let result = ListFilesResult {
repo_id: repo_id.to_owned(),
file_count: entries.len(),
files: entries,
total_bytes,
cached_count: if show_cached {
Some(cached_count)
} else {
None
},
quant_alternatives,
size_min,
size_max,
};
emit_json(&result)
}
fn parse_size_arg(s: &str) -> Result<u64, String> {
const KIB: u64 = 1024;
const MIB: u64 = 1024 * 1024;
const GIB: u64 = 1024 * 1024 * 1024;
const TIB: u64 = 1024 * GIB;
let trimmed = s.trim();
if trimmed.is_empty() {
return Err("empty size value".to_owned());
}
let split_at = trimmed
.find(|c: char| !c.is_ascii_digit() && c != '.')
.unwrap_or(trimmed.len());
let (number_part, suffix_part) = trimmed.split_at(split_at);
if number_part.is_empty() {
return Err(format!("missing number in size value {s:?}"));
}
let suffix = suffix_part.trim();
let unit: u64 = match suffix.to_ascii_lowercase().as_str() {
"" | "b" => 1,
"kib" => KIB,
"mib" => MIB,
"gib" => GIB,
"tib" => TIB,
"kb" | "mb" | "gb" | "tb" => {
return Err(format!(
"decimal size unit {suffix:?} not supported (use binary units: KiB, MiB, GiB, TiB)"
));
}
_ => {
return Err(format!(
"unrecognized size unit {suffix:?} (expected B, KiB, MiB, GiB, TiB)"
));
}
};
if !number_part.contains('.') {
let n: u64 = number_part
.parse()
.map_err(|e| format!("invalid number in size value {s:?}: {e}"))?;
return n
.checked_mul(unit)
.ok_or_else(|| format!("size value {s:?} overflows u64"));
}
let n: f64 = number_part
.parse()
.map_err(|e| format!("invalid number in size value {s:?}: {e}"))?;
if !n.is_finite() || n < 0.0 {
return Err(format!("invalid number in size value {s:?}"));
}
#[allow(clippy::cast_precision_loss, clippy::as_conversions)]
let unit_f = unit as f64;
let bytes_f = n * unit_f;
#[allow(clippy::cast_precision_loss, clippy::as_conversions)]
let max_f = u64::MAX as f64;
if !bytes_f.is_finite() || bytes_f < 0.0 || bytes_f > max_f {
return Err(format!("size value {s:?} overflows u64"));
}
#[allow(
clippy::cast_possible_truncation,
clippy::cast_sign_loss,
clippy::as_conversions
)]
let bytes = bytes_f as u64;
Ok(bytes)
}
fn format_age(time: std::time::SystemTime) -> String {
const HOUR: u64 = 3600;
const DAY: u64 = 86_400;
const MONTH: u64 = 30 * DAY;
const YEAR: u64 = 365 * DAY;
let Ok(elapsed) = time.elapsed() else {
return "\u{2014}".to_owned(); };
let secs = elapsed.as_secs();
if secs < HOUR {
"< 1 hour".to_owned()
} else if secs < DAY {
let hours = secs / HOUR;
if hours == 1 {
"1 hour ago".to_owned()
} else {
format!("{hours} hours ago")
}
} else if secs < MONTH {
let days = secs / DAY;
if days == 1 {
"1 day ago".to_owned()
} else {
format!("{days} days ago")
}
} else if secs < YEAR {
let months = secs / MONTH;
if months == 1 {
"1 month ago".to_owned()
} else {
format!("{months} months ago")
}
} else {
let years = secs / YEAR;
if years == 1 {
"1 year ago".to_owned()
} else {
format!("{years} years ago")
}
}
}
fn format_short_age(age: std::time::Duration) -> String {
const MINUTE: u64 = 60;
const HOUR: u64 = 60 * MINUTE;
const DAY: u64 = 24 * HOUR;
let secs = age.as_secs();
if secs < MINUTE {
format!("{secs}s")
} else if secs < HOUR {
format!("{}m", secs / MINUTE)
} else if secs < DAY {
format!("{}h", secs / HOUR)
} else {
format!("{}d", secs / DAY)
}
}
fn resolve_flat_target(output_dir: Option<&Path>) -> Result<PathBuf, FetchError> {
match output_dir {
Some(dir) => Ok(dir.to_path_buf()),
None => std::env::current_dir().map_err(|e| FetchError::Io {
path: PathBuf::from("."),
source: e,
}),
}
}
fn flatten_files(
file_map: &HashMap<String, PathBuf>,
target_dir: &Path,
) -> Result<Vec<PathBuf>, FetchError> {
std::fs::create_dir_all(target_dir).map_err(|e| FetchError::Io {
path: target_dir.to_path_buf(),
source: e,
})?;
let mut flat_paths = Vec::with_capacity(file_map.len());
for (filename, cache_path) in file_map {
let basename = Path::new(filename)
.file_name()
.unwrap_or(std::ffi::OsStr::new(filename.as_str()));
let flat_path = target_dir.join(basename);
std::fs::copy(cache_path, &flat_path).map_err(|e| FetchError::Io {
path: flat_path.clone(),
source: e,
})?;
flat_paths.push(flat_path);
}
Ok(flat_paths)
}
fn flatten_single_file(cache_path: &Path, target_dir: &Path) -> Result<PathBuf, FetchError> {
std::fs::create_dir_all(target_dir).map_err(|e| FetchError::Io {
path: target_dir.to_path_buf(),
source: e,
})?;
let basename = cache_path
.file_name()
.unwrap_or(std::ffi::OsStr::new("file"));
let flat_path = target_dir.join(basename);
std::fs::copy(cache_path, &flat_path).map_err(|e| FetchError::Io {
path: flat_path.clone(),
source: e,
})?;
Ok(flat_path)
}
fn warn_redundant_filters(preset: &Preset, filters: &[String]) {
let (preset_globs, preset_name): (&[&str], &str) = match preset {
Preset::Safetensors => (&["*.safetensors", "*.json", "*.txt"], "safetensors"),
Preset::Gguf => (&["*.gguf", "*.json", "*.txt"], "gguf"),
Preset::Npz => {
let globs: &[&str] = &["*.npz", "*.npy", "config.yaml", "*.json", "*.txt"];
(globs, "npz")
}
Preset::Pth => {
let globs: &[&str] = &["pytorch_model*.bin", "*.json", "*.txt"];
(globs, "pth")
}
Preset::ConfigOnly => (&["*.json", "*.txt", "*.md"], "config-only"),
};
for filter in filters {
if preset_globs.contains(&filter.as_str()) {
eprintln!("warning: --filter \"{filter}\" is redundant with --preset {preset_name}");
}
}
}
fn walk_dir_size(dir: &Path) -> u64 {
let Ok(entries) = std::fs::read_dir(dir) else {
return 0;
};
let mut total: u64 = 0;
for entry in entries.flatten() {
let Ok(meta) = entry.metadata() else {
continue;
};
if meta.is_dir() {
total = total.saturating_add(walk_dir_size(&entry.path()));
} else {
total = total.saturating_add(meta.len());
}
}
total
}
fn print_download_summary(path: &Path, elapsed: Duration) {
let total_bytes = if path.is_dir() {
walk_dir_size(path)
} else {
std::fs::metadata(path).map_or(0, |m| m.len())
};
let elapsed_secs = elapsed.as_secs_f64();
if total_bytes > 0 && elapsed_secs > 0.0 {
#[allow(clippy::cast_precision_loss, clippy::as_conversions)]
let throughput = total_bytes as f64 / elapsed_secs / (1024.0 * 1024.0);
println!(
" {} in {:.1}s ({:.1} MiB/s)",
format_size(total_bytes),
elapsed_secs,
throughput
);
}
}
fn format_downloads(n: u64) -> String {
let s = n.to_string();
let mut result = String::with_capacity(s.len() + s.len() / 3);
for (i, ch) in s.chars().enumerate() {
if i > 0 && (s.len() - i).is_multiple_of(3) {
result.push(',');
}
result.push(ch);
}
result
}
#[cfg(test)]
#[allow(
clippy::panic,
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing
)]
mod tests {
use super::*;
#[test]
fn auth_status_error_detects_401_and_403_http_errors() {
let forbidden = FetchError::Http(
"Range request for model-00002-of-00004.safetensors returned status 403 Forbidden"
.to_owned(),
);
let unauthorized = FetchError::Http(
"shard index request for org/model returned status 401 Unauthorized".to_owned(),
);
assert!(is_auth_status_error(&forbidden));
assert!(is_auth_status_error(&unauthorized));
}
#[test]
fn matches_filter_is_case_insensitive() {
assert!(matches_filter(
"model.layers.0.mlp.down_proj.weight",
"layers.0"
));
assert!(matches_filter(
"model.layers.0.mlp.down_proj.weight",
"Layers.0"
));
assert!(matches_filter(
"model.layers.0.mlp.down_proj.weight",
"LAYERS"
));
assert!(matches_filter("BLK.0.ATTN_Q.weight", "blk.0"));
assert!(!matches_filter("model.embed_tokens.weight", "layers.0"));
assert!(matches_filter("anything", ""));
}
#[test]
fn auth_status_error_ignores_other_errors() {
let not_found = FetchError::Http(
"Range request for x.safetensors returned status 404 Not Found".to_owned(),
);
let io_like = FetchError::InvalidArgument("403 in a filename is not a status".to_owned());
assert!(!is_auth_status_error(¬_found));
assert!(!is_auth_status_error(&io_like));
}
#[test]
fn enrich_gated_content_error_passes_non_auth_errors_through() {
let original =
FetchError::Http("Range request for x returned status 404 Not Found".to_owned());
let enriched = enrich_gated_content_error(original, "org/model", None);
assert!(
matches!(&enriched, FetchError::Http(msg) if msg.contains("404")),
"expected the original Http error to pass through, got: {enriched}"
);
}
fn sample_listing() -> Vec<(String, u64)> {
vec![
("model-00001-of-00002.safetensors".to_owned(), 100),
("model-00002-of-00002.safetensors".to_owned(), 200),
("params.npz".to_owned(), 50),
(
"transformer/demonCORESFWNSFW_fluxV13.safetensors".to_owned(),
300,
),
]
}
#[test]
fn narrow_pick_candidates_no_needle_keeps_all() {
let entries = sample_listing();
assert_eq!(narrow_pick_candidates(&entries, None).len(), entries.len());
}
#[test]
fn narrow_pick_candidates_substring_is_case_insensitive() {
let entries = sample_listing();
let hits = narrow_pick_candidates(&entries, Some("demoncore"));
assert_eq!(hits.len(), 1);
assert_eq!(
hits[0].0,
"transformer/demonCORESFWNSFW_fluxV13.safetensors"
);
}
#[test]
fn narrow_pick_candidates_matches_path_prefix_too() {
let entries = sample_listing();
let hits = narrow_pick_candidates(&entries, Some("TRANSFORMER/"));
assert_eq!(hits.len(), 1);
}
#[test]
fn narrow_pick_candidates_multiple_matches_preserve_order() {
let entries = sample_listing();
let hits = narrow_pick_candidates(&entries, Some("model-0000"));
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].0, "model-00001-of-00002.safetensors");
assert_eq!(hits[1].0, "model-00002-of-00002.safetensors");
}
#[test]
fn narrow_pick_candidates_no_match_is_empty() {
let entries = sample_listing();
assert!(narrow_pick_candidates(&entries, Some("gguf")).is_empty());
}
#[test]
fn parse_pick_input_accepts_in_range_with_whitespace() {
assert_eq!(parse_pick_input("2", 3), Some(2));
assert_eq!(parse_pick_input(" 3 \n", 3), Some(3));
assert_eq!(parse_pick_input("1", 1), Some(1));
}
#[test]
fn parse_pick_input_rejects_zero_out_of_range_and_garbage() {
assert_eq!(parse_pick_input("0", 3), None);
assert_eq!(parse_pick_input("4", 3), None);
assert_eq!(parse_pick_input("foo", 3), None);
assert_eq!(parse_pick_input("-1", 3), None);
assert_eq!(parse_pick_input("2.5", 3), None);
assert_eq!(parse_pick_input("", 3), None);
}
#[test]
fn multi_file_source_label_empty_is_unknown() {
assert_eq!(multi_file_source_label(&[]), "unknown");
}
#[test]
fn multi_file_source_label_all_cached() {
use inspect::InspectSource::Cached;
assert_eq!(multi_file_source_label(&[Cached, Cached]), "cached");
}
#[test]
fn multi_file_source_label_all_remote() {
use inspect::InspectSource::Remote;
assert_eq!(multi_file_source_label(&[Remote, Remote, Remote]), "remote");
}
#[test]
fn multi_file_source_label_mixed_reports_counts() {
use inspect::InspectSource::{Cached, Remote};
assert_eq!(
multi_file_source_label(&[Cached, Remote, Remote]),
"mixed (1 cached, 2 remote)"
);
}
#[test]
fn multi_file_source_label_single_cached() {
use inspect::InspectSource::Cached;
assert_eq!(multi_file_source_label(&[Cached]), "cached");
}
#[test]
fn parse_size_arg_plain_integer() {
assert_eq!(parse_size_arg("1024").unwrap(), 1024);
assert_eq!(parse_size_arg("0").unwrap(), 0);
}
#[test]
fn parse_size_arg_bytes_suffix() {
assert_eq!(parse_size_arg("512B").unwrap(), 512);
assert_eq!(parse_size_arg("512b").unwrap(), 512);
}
#[test]
fn parse_size_arg_binary_suffixes() {
assert_eq!(parse_size_arg("1KiB").unwrap(), 1024);
assert_eq!(parse_size_arg("1MiB").unwrap(), 1024 * 1024);
assert_eq!(parse_size_arg("1GiB").unwrap(), 1024 * 1024 * 1024);
assert_eq!(parse_size_arg("1TiB").unwrap(), 1024_u64.pow(4));
}
#[test]
fn parse_size_arg_case_insensitive() {
assert_eq!(parse_size_arg("5gib").unwrap(), 5 * 1024 * 1024 * 1024);
assert_eq!(parse_size_arg("5GIB").unwrap(), 5 * 1024 * 1024 * 1024);
assert_eq!(parse_size_arg("5GiB").unwrap(), 5 * 1024 * 1024 * 1024);
}
#[test]
fn parse_size_arg_whitespace_tolerant() {
assert_eq!(parse_size_arg(" 5GiB ").unwrap(), 5 * 1024 * 1024 * 1024);
assert_eq!(parse_size_arg("5 GiB").unwrap(), 5 * 1024 * 1024 * 1024);
}
#[test]
fn parse_size_arg_fractional() {
assert_eq!(
parse_size_arg("1.5GiB").unwrap(),
1024 * 1024 * 1024 * 3 / 2
);
assert_eq!(parse_size_arg("0.5KiB").unwrap(), 512);
assert_eq!(parse_size_arg(".5MiB").unwrap(), 512 * 1024);
}
#[test]
fn parse_size_arg_rejects_decimal_suffix() {
for input in ["5KB", "5MB", "5GB", "5TB", "5gb"] {
let err = parse_size_arg(input).unwrap_err();
assert!(
err.contains("decimal size unit") && err.contains("binary units"),
"input {input:?} should be rejected as decimal, got: {err}"
);
}
}
#[test]
fn parse_size_arg_rejects_unknown_suffix() {
let err = parse_size_arg("5xyz").unwrap_err();
assert!(err.contains("unrecognized size unit"), "got: {err}");
}
#[test]
fn parse_size_arg_rejects_empty() {
assert!(parse_size_arg("").unwrap_err().contains("empty"));
assert!(parse_size_arg(" ").unwrap_err().contains("empty"));
}
#[test]
fn parse_size_arg_rejects_missing_digits() {
assert!(
parse_size_arg("GiB")
.unwrap_err()
.contains("missing number")
);
assert!(
parse_size_arg("-5GiB")
.unwrap_err()
.contains("missing number")
);
}
#[test]
fn parse_size_arg_rejects_malformed_number() {
assert!(parse_size_arg("1.2.3GiB").is_err());
}
#[test]
fn parse_size_arg_overflow() {
let huge = format!("{}KiB", u64::MAX);
assert!(
parse_size_arg(huge.as_str())
.unwrap_err()
.contains("overflow")
);
}
use std::time::{Duration, SystemTime};
fn make_summary(
repo_id: &str,
size: u64,
mtime: Option<SystemTime>,
has_partial: bool,
) -> cache::CachedModelSummary {
cache::CachedModelSummary {
repo_id: repo_id.to_owned(),
file_count: 1,
total_size: size,
has_partial,
last_modified: mtime,
gguf_size_range: None,
}
}
fn fixed_now() -> SystemTime {
SystemTime::UNIX_EPOCH + Duration::from_secs(1_700_000_000)
}
#[test]
fn select_age_evictions_skips_none_mtime() {
let summaries = [make_summary("a/b", 100, None, false)];
let set = select_age_evictions(&summaries, 86_400, fixed_now());
assert!(set.is_empty(), "None mtime must not be age-evicted");
}
#[test]
fn select_age_evictions_skips_future_mtime() {
let now = fixed_now();
let future = now + Duration::from_secs(86_400);
let summaries = [make_summary("a/b", 100, Some(future), false)];
let set = select_age_evictions(&summaries, 0, now);
assert!(
set.is_empty(),
"future-dated mtime should be treated as age 0, not selected"
);
}
#[test]
fn select_age_evictions_includes_older_excludes_recent() {
let now = fixed_now();
let old = now - Duration::from_secs(86_400 * 31);
let recent = now - Duration::from_secs(86_400 * 5);
let summaries = [
make_summary("old/repo", 100, Some(old), false),
make_summary("recent/repo", 100, Some(recent), false),
];
let set = select_age_evictions(&summaries, 86_400 * 30, now);
assert!(set.contains("old/repo"));
assert!(!set.contains("recent/repo"));
}
#[test]
fn select_age_evictions_zero_threshold_evicts_everything_with_known_mtime() {
let now = fixed_now();
let summaries = [
make_summary("a/b", 100, Some(now - Duration::from_secs(1)), false),
make_summary("c/d", 100, None, false),
];
let set = select_age_evictions(&summaries, 0, now);
assert!(set.contains("a/b"));
assert!(
!set.contains("c/d"),
"None mtime still skipped at threshold 0"
);
}
fn empty_criteria() -> GcCriteria {
GcCriteria {
older_than_secs: None,
max_size: None,
except: HashSet::new(),
}
}
#[test]
fn compute_gc_plan_age_only() {
let now = fixed_now();
let summaries = [
make_summary(
"old/a",
1_000,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary(
"new/b",
2_000,
Some(now - Duration::from_secs(86_400 * 5)),
false,
),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 1);
assert_eq!(plan.evict[0].repo_id, "old/a");
assert_eq!(plan.size_before, 3_000);
assert_eq!(plan.size_after, 2_000);
assert!(!plan.budget_shortfall);
}
#[test]
fn compute_gc_plan_size_only_oldest_first() {
let now = fixed_now();
let summaries = [
make_summary(
"newest/c",
500,
Some(now - Duration::from_secs(86_400)),
false,
),
make_summary(
"oldest/a",
400,
Some(now - Duration::from_secs(86_400 * 30)),
false,
),
make_summary(
"middle/b",
300,
Some(now - Duration::from_secs(86_400 * 10)),
false,
),
];
let mut crit = empty_criteria();
crit.max_size = Some(700); let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 2);
assert_eq!(plan.evict[0].repo_id, "oldest/a");
assert_eq!(plan.evict[1].repo_id, "middle/b");
assert_eq!(plan.size_after, 500);
assert!(!plan.budget_shortfall);
}
#[test]
fn compute_gc_plan_combined_age_first_then_budget() {
let now = fixed_now();
let summaries = [
make_summary(
"ancient/a",
100,
Some(now - Duration::from_secs(86_400 * 90)),
false,
),
make_summary(
"old/b",
100,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary(
"midage/c",
400,
Some(now - Duration::from_secs(86_400 * 20)),
false,
),
make_summary(
"fresh/d",
400,
Some(now - Duration::from_secs(86_400)),
false,
),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30); crit.max_size = Some(500); let plan = compute_gc_plan(&summaries, &crit, now, false);
let evicted_ids: Vec<&str> = plan.evict.iter().map(|e| e.repo_id.as_str()).collect();
assert_eq!(evicted_ids, vec!["ancient/a", "old/b", "midage/c"]);
assert_eq!(plan.size_after, 400);
assert!(!plan.budget_shortfall);
}
#[test]
fn compute_gc_plan_except_protects_repo() {
let now = fixed_now();
let summaries = [
make_summary(
"old/a",
500,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary(
"old/b",
500,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
crit.except = HashSet::from(["old/a".to_owned()]);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 1);
assert_eq!(plan.evict[0].repo_id, "old/b");
assert_eq!(plan.protected.len(), 1);
assert_eq!(plan.protected[0].repo_id, "old/a");
}
#[test]
fn compute_gc_plan_except_causes_budget_shortfall() {
let now = fixed_now();
let summaries = [
make_summary(
"huge/a",
1_000,
Some(now - Duration::from_secs(86_400)),
false,
),
make_summary(
"tiny/b",
10,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
];
let mut crit = empty_criteria();
crit.max_size = Some(500); crit.except = HashSet::from(["huge/a".to_owned()]);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 1);
assert!(plan.budget_shortfall);
}
#[test]
fn compute_gc_plan_skips_fresh_partial() {
let now = fixed_now();
let summaries = [make_summary(
"active/repo",
1_000,
Some(now - Duration::from_secs(60 * 30)), true,
)];
let mut crit = empty_criteria();
crit.max_size = Some(0);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert!(plan.evict.is_empty());
assert_eq!(plan.skipped_partials.len(), 1);
assert!(plan.budget_shortfall);
}
#[test]
fn compute_gc_plan_evicts_stale_partial() {
let now = fixed_now();
let summaries = [make_summary(
"stale/repo",
1_000,
Some(now - Duration::from_secs(86_400 * 60)),
true, )];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 1);
assert!(plan.skipped_partials.is_empty());
}
#[test]
fn compute_gc_plan_deterministic_order_on_tied_mtime() {
let now = fixed_now();
let mtime = Some(now - Duration::from_secs(86_400 * 60));
let summaries = [
make_summary("zzz/last", 100, mtime, false),
make_summary("aaa/first", 100, mtime, false),
make_summary("mmm/middle", 100, mtime, false),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, false);
let order: Vec<&str> = plan.evict.iter().map(|e| e.repo_id.as_str()).collect();
assert_eq!(order, vec!["aaa/first", "mmm/middle", "zzz/last"]);
}
#[test]
fn compute_gc_plan_lists_kept_when_flag_set() {
let now = fixed_now();
let summaries = [
make_summary(
"old/a",
100,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary("new/b", 100, Some(now - Duration::from_secs(86_400)), false),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, true);
assert_eq!(plan.kept.len(), 1);
assert_eq!(plan.kept[0].repo_id, "new/b");
}
#[test]
fn compute_gc_plan_omits_kept_by_default() {
let now = fixed_now();
let summaries = [
make_summary(
"old/a",
100,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary("new/b", 100, Some(now - Duration::from_secs(86_400)), false),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert!(
plan.kept.is_empty(),
"kept list should be empty when list_kept=false"
);
}
#[test]
fn compute_gc_plan_zero_byte_repo_handled() {
let now = fixed_now();
let summaries = [
make_summary(
"empty/a",
0,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
make_summary(
"real/b",
100,
Some(now - Duration::from_secs(86_400 * 60)),
false,
),
];
let mut crit = empty_criteria();
crit.older_than_secs = Some(86_400 * 30);
let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 2);
assert_eq!(plan.size_before, 100);
assert_eq!(plan.size_after, 0);
}
#[test]
fn compute_gc_plan_unknown_mtime_sorted_first_for_budget() {
let now = fixed_now();
let summaries = [
make_summary(
"known/a",
400,
Some(now - Duration::from_secs(86_400 * 30)),
false,
),
make_summary("unknown/b", 400, None, false),
];
let mut crit = empty_criteria();
crit.max_size = Some(500); let plan = compute_gc_plan(&summaries, &crit, now, false);
assert_eq!(plan.evict.len(), 1);
assert_eq!(plan.evict[0].repo_id, "unknown/b");
}
fn make_tensor_info(
name: &str,
dtype: &str,
shape: Vec<usize>,
byte_len: u64,
) -> inspect::TensorInfo {
inspect::TensorInfo {
name: name.to_owned(),
dtype: dtype.to_owned(),
shape,
data_offsets: (0, byte_len),
}
}
#[test]
fn compute_dtype_groups_generic_works_with_hashmap_values() {
let mut map: HashMap<String, inspect::TensorInfo> = HashMap::new();
map.insert(
"a".to_owned(),
make_tensor_info("a", "BF16", vec![100, 100], 20_000),
);
map.insert(
"b".to_owned(),
make_tensor_info("b", "BF16", vec![50, 50], 5_000),
);
map.insert(
"c".to_owned(),
make_tensor_info("c", "F32", vec![10, 10], 400),
);
let vec_view: Vec<inspect::TensorInfo> = map.values().cloned().collect();
let from_slice = compute_dtype_groups(&vec_view);
let from_iter = compute_dtype_groups(map.values());
assert_eq!(from_slice.len(), from_iter.len());
for (dt, cnt, params, bytes) in &from_slice {
let row = from_iter
.iter()
.find(|r| r.0 == *dt)
.expect("dtype present in iter-based aggregation");
assert_eq!(row.1, *cnt);
assert_eq!(row.2, *params);
assert_eq!(row.3, *bytes);
}
}
#[test]
fn group_by_rollup_buckets_matched_and_other() {
let tensors = vec![
make_tensor_info("blk.0.ffn_gate_exps.weight", "Q4_K", vec![256, 2048], 3_000),
make_tensor_info("blk.1.ffn_gate_exps.weight", "Q4_K", vec![256, 2048], 3_000),
make_tensor_info("blk.0.attn_output.weight", "F16", vec![2048, 2048], 1_000),
];
let matcher =
compile_group_by_pattern("blk.*.ffn_*_exps.weight").expect("pattern compiles");
let rollup = compute_group_by_rollup(&tensors, &matcher);
assert_eq!(rollup.matched_tensors, 2);
assert_eq!(rollup.matched_bytes, 6_000);
assert_eq!(rollup.other_tensors, 1);
assert_eq!(rollup.other_bytes, 1_000);
assert_eq!(rollup.total_tensors, 3);
assert_eq!(rollup.total_bytes, 7_000);
}
#[test]
fn group_by_rollup_computes_per_layer_average_when_unambiguous() {
let tensors = vec![
make_tensor_info("blk.0.ffn_gate_exps.weight", "Q4_K", vec![256, 2048], 1_000),
make_tensor_info("blk.1.ffn_gate_exps.weight", "Q4_K", vec![256, 2048], 1_500),
make_tensor_info("blk.2.ffn_gate_exps.weight", "Q4_K", vec![256, 2048], 500),
];
let matcher =
compile_group_by_pattern("blk.*.ffn_*_exps.weight").expect("pattern compiles");
let rollup = compute_group_by_rollup(&tensors, &matcher);
assert_eq!(rollup.layer_count, Some(3));
assert_eq!(rollup.per_layer_bytes, Some(1_000)); }
#[test]
fn group_by_rollup_omits_per_layer_average_when_no_layer_matched() {
let tensors = vec![make_tensor_info(
"model.embed_tokens.weight",
"F16",
vec![100, 100],
1_000,
)];
let matcher =
compile_group_by_pattern("blk.*.ffn_*_exps.weight").expect("pattern compiles");
let rollup = compute_group_by_rollup(&tensors, &matcher);
assert_eq!(rollup.matched_tensors, 0);
assert_eq!(rollup.layer_count, None);
assert_eq!(rollup.per_layer_bytes, None);
}
#[test]
fn distinct_layer_indices_returns_none_when_two_positions_vary() {
let names = [
"blk.0.ffn_gate_exps.0.weight",
"blk.1.ffn_gate_exps.1.weight",
];
assert_eq!(distinct_layer_indices(&names), None);
}
#[test]
fn distinct_layer_indices_returns_none_when_nothing_varies() {
let names = ["blk.0.ffn_gate_exps.weight"];
assert_eq!(distinct_layer_indices(&names), None);
}
#[test]
fn distinct_layer_indices_finds_the_single_varying_position() {
let names = [
"blk.3.ffn_gate_exps.weight",
"blk.7.ffn_gate_exps.weight",
"blk.1.ffn_gate_exps.weight",
];
assert_eq!(distinct_layer_indices(&names), Some(vec![1, 3, 7]));
}
#[test]
fn compile_group_by_pattern_rejects_invalid_glob() {
let err = compile_group_by_pattern("blk.[.weight").expect_err("malformed glob rejected");
assert!(matches!(err, FetchError::InvalidPattern { .. }));
}
#[test]
fn format_short_age_buckets_by_unit() {
assert_eq!(format_short_age(Duration::from_secs(5)), "5s");
assert_eq!(format_short_age(Duration::from_secs(59)), "59s");
assert_eq!(format_short_age(Duration::from_secs(60)), "1m");
assert_eq!(format_short_age(Duration::from_secs(150)), "2m");
assert_eq!(format_short_age(Duration::from_secs(3600)), "1h");
assert_eq!(format_short_age(Duration::from_secs(7200)), "2h");
assert_eq!(format_short_age(Duration::from_secs(86_400)), "1d");
assert_eq!(format_short_age(Duration::from_secs(2 * 86_400)), "2d");
}
#[test]
fn moe_expert_pattern_compiles() {
assert!(compile_group_by_pattern(MOE_EXPERT_PATTERN).is_ok());
}
#[test]
fn bits_for_artifact_recognizes_specific_gguf_scheme_over_generic_prefix() {
assert_eq!(bits_for_artifact("Laguna-XS-2.1-Q4_K_M.gguf"), Some(4.85));
assert_eq!(bits_for_artifact("Laguna-XS-2.1.i1-IQ3_XS.gguf"), Some(3.3));
assert_eq!(bits_for_artifact("model-Q8_0.gguf"), Some(8.5));
}
#[test]
fn bits_for_artifact_recognizes_safetensors_suffixes() {
assert_eq!(bits_for_artifact("Laguna-XS-2.1-NVFP4"), Some(4.0));
assert_eq!(bits_for_artifact("Laguna-XS-2.1-FP8"), Some(8.0));
assert_eq!(bits_for_artifact("Laguna-XS-2.1-INT4"), Some(4.0));
}
#[test]
fn bits_for_artifact_returns_none_for_unknown_scheme() {
assert_eq!(bits_for_artifact("Laguna-XS-2.1-custom-quant-v9"), None);
}
fn make_repo_file(filename: &str, size: u64) -> repo::RepoFile {
repo::RepoFile {
filename: filename.to_owned(),
size: Some(size),
sha256: None,
}
}
#[test]
fn build_quant_rows_emits_one_row_per_gguf_file() {
let candidates = vec![discover::QuantCandidate::new(
"bartowski/Laguna-XS-2.1-GGUF".to_owned(),
discover::QuantVerification::Verified,
vec![
make_repo_file("Laguna-XS-2.1-Q4_K_M.gguf", 20_000),
make_repo_file("Laguna-XS-2.1-Q3_K_S.gguf", 14_000),
make_repo_file("config.json", 500),
],
)];
let rows = build_quant_rows(candidates);
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].artifact, "Laguna-XS-2.1-Q3_K_S.gguf"); assert_eq!(rows[0].size, 14_000); assert!(rows[0].is_gguf); assert_eq!(rows[1].artifact, "Laguna-XS-2.1-Q4_K_M.gguf"); }
#[test]
fn build_quant_rows_aggregates_safetensors_only_candidate_into_one_row() {
let candidates = vec![discover::QuantCandidate::new(
"poolside/Laguna-XS-2.1-NVFP4".to_owned(),
discover::QuantVerification::Unverified,
vec![
make_repo_file("model-00001-of-00002.safetensors", 10_000),
make_repo_file("model-00002-of-00002.safetensors", 10_000),
make_repo_file("tokenizer.json", 100),
],
)];
let rows = build_quant_rows(candidates);
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].artifact, "Laguna-XS-2.1-NVFP4"); assert_eq!(rows[0].size, 20_000); assert!(!rows[0].is_gguf); assert_eq!(rows[0].bits, Some(4.0)); }
#[test]
fn build_quant_rows_skips_candidate_with_no_recognized_weight_files() {
let candidates = vec![discover::QuantCandidate::new(
"someone/not-actually-a-quant".to_owned(),
discover::QuantVerification::Unverified,
vec![make_repo_file("README.md", 100)],
)];
assert!(build_quant_rows(candidates).is_empty());
}
#[test]
fn compute_offload_plan_fits_after_partial_offload() {
let verdict = compute_offload_plan(1_000, 800, 8, 100, 550);
match verdict {
FitVerdict::Offload {
n_cpu_moe,
moved_bytes,
resident_bytes,
} => {
assert_eq!(n_cpu_moe, 5); assert_eq!(moved_bytes, 500);
assert_eq!(resident_bytes, 500);
}
FitVerdict::FullGpu => panic!("expected Offload, got FullGpu"),
FitVerdict::DoesNotFit { reason } => {
panic!("expected Offload, got DoesNotFit({reason})")
}
}
}
#[test]
fn compute_offload_plan_rounds_shortfall_up_to_next_whole_layer() {
let verdict = compute_offload_plan(1_000, 900, 10, 90, 555);
match verdict {
FitVerdict::Offload {
n_cpu_moe,
moved_bytes,
..
} => {
assert_eq!(n_cpu_moe, 5);
assert_eq!(moved_bytes, 450);
}
FitVerdict::FullGpu => panic!("expected Offload, got FullGpu"),
FitVerdict::DoesNotFit { reason } => {
panic!("expected Offload, got DoesNotFit({reason})")
}
}
}
#[test]
fn compute_offload_plan_clamps_to_layer_count_and_reports_does_not_fit() {
let verdict = compute_offload_plan(20, 4, 4, 1, 10);
match verdict {
FitVerdict::DoesNotFit { reason } => {
assert!(reason.contains("full expert offload"), "got: {reason}");
}
FitVerdict::FullGpu => panic!("expected DoesNotFit, got FullGpu"),
FitVerdict::Offload { n_cpu_moe, .. } => {
panic!("expected DoesNotFit, got Offload(n_cpu_moe={n_cpu_moe})")
}
}
}
#[tokio::test]
async fn compute_fit_verdicts_concurrent_preserves_row_order() {
let rows: Vec<QuantArtifactRow> = (0..16)
.map(|i| QuantArtifactRow {
artifact: format!("artifact-{i}.safetensors"),
size: if i % 2 == 0 { 100 } else { 900 },
repo: "someone/repo".to_owned(),
is_gguf: false,
bits: None,
verification: discover::QuantVerification::Unverified,
})
.collect();
let verdicts = compute_fit_verdicts_concurrent(&rows, 500, 0, None).await;
assert_eq!(verdicts.len(), rows.len());
for (i, verdict) in verdicts.iter().enumerate() {
if i % 2 == 0 {
assert!(
matches!(verdict, FitVerdict::FullGpu),
"row {i} (size 100, under budget): expected FullGpu, got {verdict:?}"
);
} else {
assert!(
matches!(verdict, FitVerdict::DoesNotFit { .. }),
"row {i} (size 900, non-GGUF): expected DoesNotFit, got {verdict:?}"
);
}
}
}
#[test]
fn aggregate_diff_dtypes_scaled_sibling() {
let mut a: HashMap<String, inspect::TensorInfo> = HashMap::new();
a.insert(
"w1".to_owned(),
make_tensor_info("w1", "BF16", vec![1000, 1000], 2_000_000),
);
a.insert(
"w2".to_owned(),
make_tensor_info("w2", "BF16", vec![1000, 1000], 2_000_000),
);
a.insert(
"g".to_owned(),
make_tensor_info("g", "F32", vec![1000], 4_000),
);
let mut b: HashMap<String, inspect::TensorInfo> = HashMap::new();
b.insert(
"w1".to_owned(),
make_tensor_info("w1", "BF16", vec![500, 500], 500_000),
);
b.insert(
"g".to_owned(),
make_tensor_info("g", "F32", vec![500], 2_000),
);
let (rows_a, rows_b) = aggregate_diff_dtypes(&a, &b, None);
let dtypes_a: BTreeSet<&str> = rows_a.iter().map(|r| r.dtype.as_str()).collect();
let dtypes_b: BTreeSet<&str> = rows_b.iter().map(|r| r.dtype.as_str()).collect();
assert_eq!(dtypes_a, dtypes_b);
assert!(dtypes_a.contains("BF16"));
assert!(dtypes_a.contains("F32"));
let a_bf16 = rows_a
.iter()
.find(|r| r.dtype == "BF16")
.expect("BF16 present in A");
let b_bf16 = rows_b
.iter()
.find(|r| r.dtype == "BF16")
.expect("BF16 present in B");
assert!(a_bf16.bytes > b_bf16.bytes);
assert!(a_bf16.tensors > b_bf16.tensors);
}
#[test]
fn aggregate_diff_dtypes_architectural_variant() {
let mut a: HashMap<String, inspect::TensorInfo> = HashMap::new();
a.insert(
"w".to_owned(),
make_tensor_info("w", "BF16", vec![1000, 1000], 2_000_000),
);
a.insert(
"expert".to_owned(),
make_tensor_info("expert", "F8_E4M3", vec![100, 100], 10_000),
);
let mut b: HashMap<String, inspect::TensorInfo> = HashMap::new();
b.insert(
"w".to_owned(),
make_tensor_info("w", "BF16", vec![1000, 1000], 2_000_000),
);
let (rows_a, rows_b) = aggregate_diff_dtypes(&a, &b, None);
let dtypes_a: BTreeSet<&str> = rows_a.iter().map(|r| r.dtype.as_str()).collect();
let dtypes_b: BTreeSet<&str> = rows_b.iter().map(|r| r.dtype.as_str()).collect();
assert!(dtypes_a.contains("F8_E4M3"));
assert!(!dtypes_b.contains("F8_E4M3"));
assert!(dtypes_a.contains("BF16"));
assert!(dtypes_b.contains("BF16"));
}
#[test]
fn aggregate_diff_dtypes_with_filter() {
let mut a: HashMap<String, inspect::TensorInfo> = HashMap::new();
a.insert(
"expert.w".to_owned(),
make_tensor_info("expert.w", "BF16", vec![100], 200),
);
a.insert(
"attn.q".to_owned(),
make_tensor_info("attn.q", "BF16", vec![1000], 2_000),
);
let mut b: HashMap<String, inspect::TensorInfo> = HashMap::new();
b.insert(
"expert.w".to_owned(),
make_tensor_info("expert.w", "BF16", vec![50], 100),
);
b.insert(
"attn.q".to_owned(),
make_tensor_info("attn.q", "BF16", vec![500], 1_000),
);
let (rows_a, rows_b) = aggregate_diff_dtypes(&a, &b, Some("expert"));
assert_eq!(rows_a.len(), 1);
assert_eq!(rows_b.len(), 1);
assert_eq!(rows_a[0].dtype, "BF16");
assert_eq!(rows_a[0].tensors, 1);
assert_eq!(rows_a[0].bytes, 200);
assert_eq!(rows_b[0].tensors, 1);
assert_eq!(rows_b[0].bytes, 100);
}
#[test]
fn collapse_numeric_segments_empty_string() {
assert_eq!(collapse_numeric_segments(""), "");
}
#[test]
fn collapse_numeric_segments_no_digits() {
assert_eq!(
collapse_numeric_segments("model.embed_tokens.weight"),
"model.embed_tokens.weight"
);
}
#[test]
fn collapse_numeric_segments_single_digit() {
assert_eq!(
collapse_numeric_segments("model.layers.3.mlp.weight"),
"model.layers.{N}.mlp.weight"
);
}
#[test]
fn collapse_numeric_segments_multi_digit_run() {
assert_eq!(
collapse_numeric_segments("model.layers.123.mlp.weight"),
"model.layers.{N}.mlp.weight"
);
}
#[test]
fn collapse_numeric_segments_multiple_runs() {
assert_eq!(
collapse_numeric_segments("block.24.attn.qkv.weight.blocks.7"),
"block.{N}.attn.qkv.weight.blocks.{N}"
);
}
#[test]
fn aggregate_diff_collapse_only_a_groups_and_sums_by_pattern() {
let tensors_a: HashMap<String, inspect::TensorInfo> = [
(
"model.layers.0.mlp.gate_proj.weight_scale".to_owned(),
make_tensor_info(
"model.layers.0.mlp.gate_proj.weight_scale",
"F32",
vec![1],
4,
),
),
(
"model.layers.1.mlp.gate_proj.weight_scale".to_owned(),
make_tensor_info(
"model.layers.1.mlp.gate_proj.weight_scale",
"F32",
vec![1],
4,
),
),
(
"lm_head.weight".to_owned(),
make_tensor_info("lm_head.weight", "BF16", vec![100, 100], 20_000),
),
]
.into_iter()
.collect();
let names: Vec<&str> = vec![
"model.layers.0.mlp.gate_proj.weight_scale",
"model.layers.1.mlp.gate_proj.weight_scale",
"lm_head.weight",
];
let rows = aggregate_diff_collapse(&names, &tensors_a, None);
assert_eq!(rows.len(), 2, "two distinct patterns");
let layer_row = rows
.iter()
.find(|r| r.pattern == "model.layers.{N}.mlp.gate_proj.weight_scale")
.expect("layer pattern present");
assert_eq!(layer_row.tensors, 2);
assert_eq!(layer_row.bytes, 8);
assert!(layer_row.bytes_b.is_none(), "only-A rows carry no B total");
let head_row = rows
.iter()
.find(|r| r.pattern == "lm_head.weight")
.expect("lm_head pattern present, unchanged by digit collapse");
assert_eq!(head_row.tensors, 1);
assert_eq!(head_row.bytes, 20_000);
}
#[test]
fn aggregate_diff_collapse_only_b_sums_from_b_side() {
let tensors_b: HashMap<String, inspect::TensorInfo> = [
(
"model.layers.0.mlp.down_proj.qweight".to_owned(),
make_tensor_info("model.layers.0.mlp.down_proj.qweight", "I32", vec![10], 40),
),
(
"model.layers.1.mlp.down_proj.qweight".to_owned(),
make_tensor_info("model.layers.1.mlp.down_proj.qweight", "I32", vec![10], 40),
),
]
.into_iter()
.collect();
let names: Vec<&str> = vec![
"model.layers.0.mlp.down_proj.qweight",
"model.layers.1.mlp.down_proj.qweight",
];
let rows = aggregate_diff_collapse(&names, &tensors_b, None);
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].pattern, "model.layers.{N}.mlp.down_proj.qweight");
assert_eq!(rows[0].tensors, 2);
assert_eq!(rows[0].bytes, 80);
assert!(rows[0].bytes_b.is_none());
}
#[test]
fn aggregate_diff_collapse_differ_populates_both_sides_and_sorts_by_bytes_desc() {
let tensors_a: HashMap<String, inspect::TensorInfo> = [
(
"model.layers.0.input_layernorm.weight".to_owned(),
make_tensor_info(
"model.layers.0.input_layernorm.weight",
"F16",
vec![100],
200,
),
),
(
"model.embed_tokens.weight".to_owned(),
make_tensor_info("model.embed_tokens.weight", "F16", vec![1000, 100], 200_000),
),
]
.into_iter()
.collect();
let tensors_b: HashMap<String, inspect::TensorInfo> = [
(
"model.layers.0.input_layernorm.weight".to_owned(),
make_tensor_info(
"model.layers.0.input_layernorm.weight",
"BF16",
vec![100],
200,
),
),
(
"model.embed_tokens.weight".to_owned(),
make_tensor_info(
"model.embed_tokens.weight",
"BF16",
vec![1000, 100],
200_000,
),
),
]
.into_iter()
.collect();
let names: Vec<&str> = vec![
"model.layers.0.input_layernorm.weight",
"model.embed_tokens.weight",
];
let rows = aggregate_diff_collapse(&names, &tensors_a, Some(&tensors_b));
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].pattern, "model.embed_tokens.weight");
assert_eq!(rows[0].bytes, 200_000);
assert_eq!(rows[0].bytes_b, Some(200_000));
assert_eq!(rows[1].pattern, "model.layers.{N}.input_layernorm.weight");
assert_eq!(rows[1].bytes, 200);
assert_eq!(rows[1].bytes_b, Some(200));
}
#[allow(clippy::field_reassign_with_default)] fn base_model_config() -> inspect::ModelConfig {
let mut c = inspect::ModelConfig::default();
c.model_type = Some("llama".to_owned());
c.num_hidden_layers = Some(16);
c.hidden_size = Some(2048);
c
}
#[test]
fn build_config_diff_rows_identical_configs_report_no_differences() {
let a = base_model_config();
let b = base_model_config();
let rows = build_config_diff_rows(&a, &b);
assert_eq!(rows.len(), 23, "one row per ModelConfig field");
assert_eq!(rows.iter().filter(|r| r.differs).count(), 0);
}
#[test]
fn build_config_diff_rows_detects_differing_numeric_field() {
let a = base_model_config();
let mut b = base_model_config();
b.num_hidden_layers = Some(32);
let rows = build_config_diff_rows(&a, &b);
let differing: Vec<&ConfigFieldDiff> = rows.iter().filter(|r| r.differs).collect();
assert_eq!(differing.len(), 1, "only num_hidden_layers should differ");
assert_eq!(differing[0].field, "num_hidden_layers");
assert_eq!(differing[0].a.as_deref(), Some("16"));
assert_eq!(differing[0].b.as_deref(), Some("32"));
}
#[test]
fn build_config_diff_rows_field_present_on_one_side_counts_as_differing() {
let a = base_model_config();
let mut b = base_model_config();
b.sliding_window = Some(4096);
let rows = build_config_diff_rows(&a, &b);
let row = rows
.iter()
.find(|r| r.field == "sliding_window")
.expect("sliding_window row present");
assert!(row.differs, "None vs Some(4096) must count as differing");
assert!(row.a.is_none());
assert_eq!(row.b.as_deref(), Some("4096"));
}
#[test]
fn build_config_diff_rows_vec_field_detects_a_late_difference() {
let mut a = base_model_config();
a.layer_types = Some(vec!["attention".to_owned(); 20]);
let mut b = base_model_config();
let mut layer_types_b = vec!["attention".to_owned(); 20];
#[allow(clippy::indexing_slicing)]
{
layer_types_b[19] = "mamba".to_owned();
}
b.layer_types = Some(layer_types_b);
let rows = build_config_diff_rows(&a, &b);
let row = rows
.iter()
.find(|r| r.field == "layer_types")
.expect("layer_types row present");
assert!(
row.differs,
"a difference in the 20th element must still be detected"
);
}
#[test]
fn truncate_for_display_short_string_unchanged() {
assert_eq!(truncate_for_display("llama"), "llama");
}
#[test]
fn truncate_for_display_long_string_truncates_with_ellipsis() {
let long = "a".repeat(CONFIG_CELL_DISPLAY_CAP + 10);
let shown = truncate_for_display(&long);
assert_eq!(shown.chars().count(), CONFIG_CELL_DISPLAY_CAP + 1); assert!(shown.ends_with('\u{2026}'));
}
#[test]
fn diff_tensor_side_serializes_byte_count() {
let side = DiffTensorSide {
dtype: "BF16".to_owned(),
shape: vec![10, 10],
byte_count: 200,
};
let json_str = serde_json::to_string(&side).expect("DiffTensorSide serializes cleanly");
assert!(
json_str.contains("\"byte_count\":200"),
"expected byte_count in JSON, got: {json_str}"
);
assert!(json_str.contains("\"dtype\":\"BF16\""));
assert!(json_str.contains("\"shape\":[10,10]"));
}
#[test]
fn diff_dtypes_canonical_case_fits_design_width() {
let rows_a = vec![
DiffDtypeGroup {
dtype: "BF16".to_owned(),
tensors: 630,
params: 3_610_000_000,
bytes: 7_213_120_000,
},
DiffDtypeGroup {
dtype: "U8".to_owned(),
tensors: 192,
params: 20_300_000_000,
bytes: 20_303_437_824,
},
DiffDtypeGroup {
dtype: "F32".to_owned(),
tensors: 50,
params: 70_000_000,
bytes: 322_122_547,
},
];
let rows_b = vec![
DiffDtypeGroup {
dtype: "BF16".to_owned(),
tensors: 126,
params: 720_000_000,
bytes: 1_438_986_240,
},
DiffDtypeGroup {
dtype: "U8".to_owned(),
tensors: 192,
params: 16_000_000_000,
bytes: 15_351_808_000,
},
DiffDtypeGroup {
dtype: "F32".to_owned(),
tensors: 40,
params: 40_000_000,
bytes: 268_435_456,
},
];
let w = diff_dtypes_column_widths(&rows_a, &rows_b);
let total = w.total_width();
assert!(
total <= DIFF_DTYPES_DESIGN_WIDTH,
"canonical histogram width = {total} chars, design target {DIFF_DTYPES_DESIGN_WIDTH}",
);
}
fn make_block_tensors(
start: usize,
count: usize,
outlier_indices: &[usize],
) -> Vec<inspect::TensorInfo> {
let mut tensors = Vec::new();
for i in start..start + count {
tensors.push(make_tensor_info(
&format!("blocks.{i}.x.weight"),
"F32",
vec![4],
16,
));
tensors.push(make_tensor_info(
&format!("blocks.{i}.y.weight"),
"F32",
vec![8],
32,
));
if outlier_indices.contains(&i) {
tensors.push(make_tensor_info(
&format!("blocks.{i}.z.weight"),
"F32",
vec![2],
8,
));
}
}
tensors
}
fn blocks_children(forest: &[TreeNode]) -> &[TreeNode] {
forest
.iter()
.find_map(|n| match n {
TreeNode::Branch(b) if b.segment == "blocks" => Some(b.children.as_slice()),
TreeNode::Leaf(_) | TreeNode::Branch(_) | TreeNode::Ranged(_) => None,
})
.expect("forest must contain a top-level `blocks` branch")
}
#[test]
fn tree_collapse_full_range_unchanged_with_no_outlier() {
let tensors = make_block_tensors(0, 4, &[]);
let forest = collapse_ranges(build_tree(&tensors));
assert_eq!(forest.len(), 1, "expected one top-level `blocks` node");
match &forest[0] {
TreeNode::Ranged(r) => {
assert_eq!(r.segment, "blocks");
assert_eq!((r.range_start, r.range_end), (0, 3));
assert_eq!(r.template.len(), 2, "template is one block's 2 tensors");
}
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected a bare Ranged node, got {other:?}")
}
}
}
#[test]
fn tree_collapse_tolerates_leading_outlier() {
let tensors = make_block_tensors(0, 4, &[0]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(children.len(), 2, "outlier + one collapsed range");
match &children[0] {
TreeNode::Branch(b) => {
assert_eq!(b.segment, "0");
assert_eq!(b.children.len(), 3, "block 0 has its extra z.weight tensor");
}
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block 0 rendered standalone, got {other:?}")
}
}
match &children[1] {
TreeNode::Ranged(r) => {
assert_eq!(
r.segment, "",
"nested range omits the repeated `blocks` segment"
);
assert_eq!((r.range_start, r.range_end), (1, 3));
assert_eq!(r.template.len(), 2);
}
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected the [1..3] range, got {other:?}")
}
}
}
#[test]
fn tree_collapse_tolerates_trailing_outlier() {
let tensors = make_block_tensors(0, 4, &[3]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(children.len(), 2, "one collapsed range + outlier");
match &children[0] {
TreeNode::Ranged(r) => {
assert_eq!(r.segment, "");
assert_eq!((r.range_start, r.range_end), (0, 2));
}
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected the [0..2] range, got {other:?}")
}
}
match &children[1] {
TreeNode::Branch(b) => {
assert_eq!(b.segment, "3");
assert_eq!(b.children.len(), 3);
}
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block 3 rendered standalone, got {other:?}")
}
}
}
#[test]
fn tree_collapse_tolerates_outlier_at_both_edges() {
let mut tensors = make_block_tensors(0, 5, &[0]);
tensors.push(make_tensor_info("blocks.4.w.weight", "F32", vec![1], 4));
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(
children.len(),
3,
"leading outlier + range + trailing outlier"
);
match &children[0] {
TreeNode::Branch(b) => assert_eq!(b.segment, "0"),
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block 0 standalone, got {other:?}")
}
}
match &children[1] {
TreeNode::Ranged(r) => assert_eq!((r.range_start, r.range_end), (1, 3)),
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected the [1..3] range, got {other:?}")
}
}
match &children[2] {
TreeNode::Branch(b) => assert_eq!(b.segment, "4"),
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block 4 standalone, got {other:?}")
}
}
}
#[test]
fn tree_collapse_tolerates_k_outliers_per_edge() {
let tensors = make_block_tensors(0, 7, &[0, 1, 2]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(children.len(), 4, "3 outliers + one collapsed range");
for (i, child) in children.iter().take(3).enumerate() {
match child {
TreeNode::Branch(b) => assert_eq!(b.segment, i.to_string()),
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block {i} standalone, got {other:?}")
}
}
}
#[allow(clippy::indexing_slicing)]
match &children[3] {
TreeNode::Ranged(r) => assert_eq!((r.range_start, r.range_end), (3, 6)),
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected the [3..6] range, got {other:?}")
}
}
}
#[test]
fn tree_collapse_refuses_more_than_k_outliers_at_one_edge() {
let tensors = make_block_tensors(0, 8, &[0, 1, 2, 3]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(
children.len(),
8,
"no collapse at all — every block stands alone"
);
for child in children {
assert!(
matches!(child, TreeNode::Branch(_)),
"expected every child to remain an individual Branch, got {child:?}"
);
}
}
#[test]
fn tree_collapse_tolerates_k_outliers_at_both_edges() {
let tensors = make_block_tensors(0, 9, &[0, 1, 2, 6, 7, 8]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(
children.len(),
7,
"3 leading outliers + range + 3 trailing outliers"
);
for (i, child) in children.iter().take(3).enumerate() {
match child {
TreeNode::Branch(b) => assert_eq!(b.segment, i.to_string()),
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block {i} standalone, got {other:?}")
}
}
}
#[allow(clippy::indexing_slicing)]
match &children[3] {
TreeNode::Ranged(r) => assert_eq!((r.range_start, r.range_end), (3, 5)),
other @ (TreeNode::Leaf(_) | TreeNode::Branch(_)) => {
panic!("expected the [3..5] range, got {other:?}")
}
}
for (offset, i) in (6..9).enumerate() {
#[allow(clippy::indexing_slicing)]
match &children[4 + offset] {
TreeNode::Branch(b) => assert_eq!(b.segment, i.to_string()),
other @ (TreeNode::Leaf(_) | TreeNode::Ranged(_)) => {
panic!("expected block {i} standalone, got {other:?}")
}
}
}
}
#[test]
fn tree_collapse_refuses_a_middle_outlier() {
let tensors = make_block_tensors(0, 9, &[4]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(
children.len(),
9,
"no collapse at all — every block stands alone"
);
for child in children {
assert!(
matches!(child, TreeNode::Branch(_)),
"expected every child to remain an individual Branch, got {child:?}"
);
}
}
#[test]
fn tree_collapse_refuses_when_too_few_survive() {
let tensors = make_block_tensors(0, 2, &[0]);
let forest = collapse_ranges(build_tree(&tensors));
let children = blocks_children(&forest);
assert_eq!(children.len(), 2);
for child in children {
assert!(matches!(child, TreeNode::Branch(_)));
}
}
#[test]
fn tree_collapse_partial_totals_match_full_expansion() {
let tensors = make_block_tensors(0, 6, &[0]);
let expected_tensors = tensors.len();
let expected_params: u64 = tensors.iter().map(inspect::TensorInfo::num_elements).sum();
let forest = collapse_ranges(build_tree(&tensors));
let TreeNode::Branch(blocks) = forest
.iter()
.find(|n| matches!(n, TreeNode::Branch(b) if b.segment == "blocks"))
.expect("blocks branch present")
else {
unreachable!("matched above");
};
assert_eq!(blocks.total_tensors, expected_tensors);
assert_eq!(blocks.total_params, expected_params);
let TreeNode::Ranged(ranged) = blocks
.children
.iter()
.find(|c| matches!(c, TreeNode::Ranged(_)))
.expect("one Ranged child present")
else {
unreachable!("matched above");
};
let count = ranged.range_end - ranged.range_start + 1;
assert_eq!(ranged.total_tensors, 2 * count);
}
#[test]
fn tree_children_sort_numerically_even_when_collapse_fails() {
let mut tensors = Vec::new();
for i in 0..12 {
tensors.push(make_tensor_info(
&format!("layers.{i}.mlp.weight"),
"F32",
vec![4],
16,
));
tensors.push(make_tensor_info(
&format!("layers.{i}.norm.weight"),
"F32",
vec![4],
16,
));
if i == 4 || i == 8 {
tensors.push(make_tensor_info(
&format!("layers.{i}.attn.weight"),
"F32",
vec![2],
8,
));
}
}
let forest = collapse_ranges(build_tree(&tensors));
assert_eq!(forest.len(), 1);
let TreeNode::Branch(layers) = &forest[0] else {
panic!(
"expected the layers branch to remain uncollapsed, got {:?}",
forest[0]
);
};
assert_eq!(layers.segment, "layers");
assert_eq!(
layers.children.len(),
12,
"middle outliers must never collapse — every layer stands alone"
);
let order: Vec<usize> = layers
.children
.iter()
.map(|c| {
let TreeNode::Branch(b) = c else {
panic!("expected every child to remain an individual Branch, got {c:?}");
};
b.segment
.parse::<usize>()
.expect("every child segment is numeric by construction")
})
.collect();
assert_eq!(
order,
(0..12).collect::<Vec<_>>(),
"children must render in numeric order, not lexicographic (0,1,10,11,...)"
);
}
#[test]
fn template_detector_fires_on_default_hf_card() {
let body = "\
# Model Card for Model ID
<!-- Provide a quick summary of what the model is/does. -->
## Model Details
### Model Description
<!-- Provide a longer summary of what this model is. -->
";
assert!(looks_like_default_template(body));
}
#[test]
fn template_detector_clears_a_real_readme() {
let body = "\
# Llama 3.2 1B Instruct
A 1B-parameter instruction-tuned Llama 3.2 model trained on a mixture of
public and proprietary data. Optimized for low-latency on-device inference.
## Usage
```python
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(\"meta-llama/Llama-3.2-1B\")
```
## License
This model is released under the Llama 3.2 community license.
";
assert!(!looks_like_default_template(body));
}
#[test]
fn template_detector_is_robust_to_blank_input() {
assert!(!looks_like_default_template(""));
}
#[test]
fn template_detector_fires_on_high_comment_density() {
let body = "\
# Random Custom Title
Some intro text here.
<!-- comment 1 -->
<!-- comment 2 -->
<!-- comment 3 -->
<!-- comment 4 -->
<!-- comment 5 -->
More content.
More content.
More content.
More content.
";
assert!(looks_like_default_template(body));
}
#[test]
fn format_header_line_safetensors_with_total() {
let line = format_header_line(64_768, Some(1_073_741_824), false);
assert!(line.starts_with("Header:"), "got: {line}");
assert!(
line.contains("(JSON)"),
"safetensors line should keep `(JSON)` suffix, got: {line}"
);
assert!(
line.contains("total"),
"safetensors line with file_size should include `total`, got: {line}"
);
}
#[test]
fn format_header_line_safetensors_without_total() {
let line = format_header_line(64_768, None, false);
assert!(line.starts_with("Header:"), "got: {line}");
assert!(line.contains("(JSON)"), "got: {line}");
assert!(
!line.contains("total"),
"no file_size → no `total` clause, got: {line}"
);
}
#[test]
fn format_header_line_headerless_with_total() {
let line = format_header_line(0, Some(3_705_032_704), true);
assert!(
line.starts_with("Size:"),
"headerless-format line should use `Size:` label, got: {line}"
);
assert!(
!line.contains("(JSON)"),
"headerless-format line must not have safetensors-flavoured `(JSON)`, got: {line}"
);
assert!(
!line.contains(" 0 B "),
"headerless-format line must not show the meaningless `0 B` prefix, got: {line}"
);
}
#[test]
fn format_header_line_headerless_without_total() {
let line = format_header_line(0, None, true);
assert_eq!(line, "Size: (size unknown)");
}
#[test]
fn format_metadata_lines_empty_returns_empty_vec() {
let meta: HashMap<String, String> = HashMap::new();
let lines = format_metadata_lines(&meta);
assert!(
lines.is_empty(),
"empty metadata → no lines, got: {lines:?}"
);
}
#[test]
fn format_metadata_lines_inline_under_threshold() {
let mut meta: HashMap<String, String> = HashMap::new();
meta.insert("quant_method".to_owned(), "gptq".to_owned());
meta.insert("bits".to_owned(), "4".to_owned());
meta.insert("group_size".to_owned(), "128".to_owned());
let lines = format_metadata_lines(&meta);
assert_eq!(
lines.len(),
1,
"≤ 6 keys → single inline line, got: {lines:?}"
);
let line = &lines[0];
assert!(
line.starts_with("Metadata: "),
"inline form starts with `Metadata: `, got: {line}"
);
assert!(line.contains("bits=4"), "got: {line}");
assert!(line.contains("group_size=128"), "got: {line}");
assert!(line.contains("quant_method=gptq"), "got: {line}");
let bits_pos = line.find("bits=").expect("bits=");
let group_pos = line.find("group_size=").expect("group_size=");
let quant_pos = line.find("quant_method=").expect("quant_method=");
assert!(bits_pos < group_pos, "expected alphabetical order");
assert!(group_pos < quant_pos, "expected alphabetical order");
}
#[test]
fn format_metadata_lines_tabular_over_threshold() {
let mut meta: HashMap<String, String> = HashMap::new();
meta.insert("general.architecture".to_owned(), "llama".to_owned());
meta.insert("general.name".to_owned(), "Mistral-7B".to_owned());
meta.insert("general.quantization".to_owned(), "Q4_K_M".to_owned());
meta.insert("llama.context_length".to_owned(), "32768".to_owned());
meta.insert("llama.head_count".to_owned(), "32".to_owned());
meta.insert("tokenizer.ggml.bos_token_id".to_owned(), "1".to_owned());
meta.insert("gguf.version".to_owned(), "3".to_owned());
let lines = format_metadata_lines(&meta);
assert!(lines.len() > 1, "> 6 keys → tabular block, got: {lines:?}");
assert_eq!(
lines[0], "Metadata:",
"first line is the `Metadata:` header"
);
let body = lines[1..].join("\n");
let arch_pos = body
.find("general.architecture=")
.expect("general.architecture");
let name_pos = body.find("general.name=").expect("general.name");
let gguf_pos = body.find("gguf.version=").expect("gguf.version");
let llama_pos = body
.find("llama.context_length=")
.expect("llama.context_length");
let tok_pos = body
.find("tokenizer.ggml.bos_token_id=")
.expect("tokenizer.*");
assert!(arch_pos < name_pos, "alphabetical within `general.*`");
assert!(name_pos < gguf_pos, "`general.*` precedes `gguf.*`");
assert!(gguf_pos < llama_pos, "`gguf.*` precedes `llama.*`");
assert!(llama_pos < tok_pos, "`llama.*` precedes `tokenizer.*`");
for line in &lines[1..] {
assert!(
line.starts_with(" "),
"tabular value line must be indented, got: {line:?}"
);
}
}
#[test]
fn format_metadata_lines_multiline_value_renders_as_block() {
let chat_template = "{%- for message in messages %}\n{{- '<|im_start|>' + message['role'] + '\\n' }}\n{%- endfor -%}";
let mut meta: HashMap<String, String> = HashMap::new();
meta.insert("a".to_owned(), "1".to_owned());
meta.insert("b".to_owned(), "2".to_owned());
meta.insert("c".to_owned(), "3".to_owned());
meta.insert("d".to_owned(), "4".to_owned());
meta.insert("e".to_owned(), "5".to_owned());
meta.insert("f".to_owned(), "6".to_owned());
meta.insert("z_chat_template".to_owned(), chat_template.to_owned());
let lines = format_metadata_lines(&meta);
let key_line_idx = lines
.iter()
.position(|l| l == " z_chat_template=")
.expect("multi-line key gets its own header line");
let next = lines
.get(key_line_idx + 1)
.expect("at least one continuation line");
assert!(
next.starts_with(" "),
"continuation lines have 4-space indent, got: {next:?}"
);
assert!(
next.contains("for message"),
"first continuation line carries the value, got: {next:?}"
);
}
#[test]
fn format_metadata_lines_sort_is_deterministic() {
let mut meta: HashMap<String, String> = HashMap::new();
for i in 0..10 {
meta.insert(format!("key_{i:02}"), format!("value_{i}"));
}
let first = format_metadata_lines(&meta);
let second = format_metadata_lines(&meta);
assert_eq!(first, second, "sort must be deterministic across runs");
}
#[test]
fn format_quant_lines_returns_empty_for_none() {
let lines = format_quant_lines(None);
assert!(lines.is_empty(), "None input → no lines, got: {lines:?}");
}
#[test]
fn format_quant_lines_returns_two_lines_for_some() {
let q = inspect::QuantInfo {
scheme: "Bnb4".to_owned(),
stored_bytes: 4_400_000_000, dequantized_bytes: 8_810_000_000, };
let lines = format_quant_lines(Some(&q));
assert_eq!(
lines.len(),
2,
"Some input → exactly two lines, got: {lines:?}"
);
assert!(
lines[0].starts_with("Format:"),
"first line is the Format: header, got: {}",
lines[0]
);
assert!(
lines[0].contains("Bnb4"),
"Format line carries the scheme string, got: {}",
lines[0]
);
assert!(
lines[1].starts_with("Size:"),
"second line is the Size: header, got: {}",
lines[1]
);
}
#[test]
fn format_quant_lines_size_uses_stored_arrow_dequantised() {
let q = inspect::QuantInfo {
scheme: "FineGrainedFp8".to_owned(),
stored_bytes: 4_400_000_000,
dequantized_bytes: 8_810_000_000,
};
let lines = format_quant_lines(Some(&q));
let size_line = &lines[1];
assert!(
size_line.contains(" stored -> "),
"Size line must use the `stored -> ` arrow separator, got: {size_line}"
);
assert!(
size_line.ends_with("(BF16)"),
"Size line must annotate the dequantised side with `(BF16)`, got: {size_line}"
);
}
}