Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions report/threejs/hf-vocab-sphere/app/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
from pathlib import Path

import numpy as np

from fastapi import FastAPI, HTTPException, Query
from fastapi.responses import HTMLResponse
from fastapi.staticfiles import StaticFiles
Expand All @@ -19,12 +18,13 @@
list_local_models,
load_model,
model_status,
model_vector_function,
nearest_neighbors,
search_tokens,
tokenize_text,
selected_token_rows,
selected_vectors,
token_record,
tokenize_text,
unload_model,
)
from .projections import nearest_neighbor_edges, project_vectors, projection_catalog
Expand All @@ -47,7 +47,7 @@
TokenSearchResponse,
TokenWindowResponse,
)
from .vector_math import VectorExpressionResult, alias_for_index, evaluate_vector_expressions
from .vector_math import VectorExpressionResult, alias_for_index, evaluate_vector_expressions, vector_dimension_metrics

BASE_DIR = Path(__file__).resolve().parent
INDEX_PATH = BASE_DIR / "templates" / "index.html"
Expand Down Expand Up @@ -237,7 +237,11 @@ def _prepare_projection_selection(
requested_ids.insert(0, anchor_id)
ids, base_vectors = selected_vectors(assets, requested_ids)
anchor_index = ids.index(anchor_id) if anchor_id is not None else None
resultants = evaluate_vector_expressions(base_vectors, arithmetic_expressions or [])
resultants = evaluate_vector_expressions(
base_vectors,
arithmetic_expressions or [],
model_function=lambda name, args: model_vector_function(assets, name, args),
)
if resultants:
vectors = np.vstack([base_vectors, *(item.vector[None, :] for item in resultants)])
else:
Expand All @@ -262,6 +266,7 @@ def _projection_rows(
"kind": "token",
"alias": alias_for_index(index),
"label": row["display"],
**vector_dimension_metrics(base_vectors[index]),
"expression": None,
"referenced_aliases": [],
}
Expand All @@ -283,6 +288,7 @@ def _projection_rows(
"special": False,
"present_in_tokenizer": False,
"magnitude": item.magnitude,
**vector_dimension_metrics(item.vector),
"rank": None,
"cosine_similarity": None,
"angle_deg": None,
Expand Down
Loading