Skip to content
Merged
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: 14 additions & 0 deletions workbench/_api/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,3 +46,17 @@ def user_has_model_access(user_email: str, model_name: str, state: "AppState") -

return True


def require_model_access(state: "AppState", user_email: str, model_name: str) -> None:
"""Refuse a caller who cannot use this model, before anything runs.

A route calls this while it can still fail the ordinary way: once it starts
streaming, the status line is gone and a refusal can only be an `error` frame
(see ``sse``). Local deployments gate nothing -- there is no catalog to be
outside of.
"""
if state.remote and not user_has_model_access(user_email, model_name, state):
raise HTTPException(
status_code=403, detail=f"User does not have access to {model_name}"
)

55 changes: 13 additions & 42 deletions workbench/_api/routes/activation_patching.py
Original file line number Diff line number Diff line change
@@ -1,19 +1,17 @@
from fastapi import APIRouter, Request, Depends
from typing import List, Union
from pydantic import BaseModel

from ..data_models import NDIFResponse
from fastapi import APIRouter, Depends
from pydantic import BaseModel

from nnsightful.tools.activation_patching import activation_patching

from ..state import AppState
from ..auth import require_user_email
from ..state import get_state

from nnsightful.types import ActivationPatchingData
from nnsightful.tools.activation_patching import activation_patching
from ..sse import stream_tool
from ..state import AppState, get_state

router = APIRouter()


class ActivationPatchingRequest(BaseModel):
model_name: str
src_prompt: str
Expand All @@ -23,48 +21,21 @@ class ActivationPatchingRequest(BaseModel):
tgt_freeze: List[int] = []
token_ids: List[int]

class ActivationPatchingResponse(NDIFResponse):
data: ActivationPatchingData | None = None


@router.post("/start", response_model=ActivationPatchingResponse)
async def start_activation_patching(
@router.post("/run")
async def run_activation_patching(
request: ActivationPatchingRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
model = state[request.model_name]
backend = state.make_backend(model=model)

output = activation_patching._run(
model,
"""Run activation patching, streaming status until the data lands (see ``sse``)."""
return stream_tool(
state,
activation_patching,
state[request.model_name],
request.src_prompt,
request.tgt_prompt,
request.src_pos,
request.tgt_pos,
request.tgt_freeze,
remote=state.remote,
backend=backend,
non_blocking=state.remote,
raw=False,
)

if not backend.blocking:
return {"job_id": output}

return {"data": activation_patching.to_data_obj(**output)}


@router.post("/results/{job_id}", response_model=ActivationPatchingResponse)
async def collect_results(
job_id: str,
request: ActivationPatchingRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
backend = state.make_backend(job_id=job_id)
results = backend()['results']

data = activation_patching.to_data_obj(**results)

return {"data": data}
82 changes: 23 additions & 59 deletions workbench/_api/routes/causal_mediation.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,9 @@
from pydantic import BaseModel, Field

from ..auth import require_user_email
from ..data_models import NDIFResponse
from ..sse import stream
from ..state import AppState, get_state

from nnsightful.types import LogitLensData

router = APIRouter()


Expand All @@ -27,12 +25,6 @@ class CausalMediationRequest(BaseModel):
include_entropy: bool = True


class CausalMediationResponse(NDIFResponse):
"""Identical shape to LogitLensResponse so the frontend can reuse the
existing logit-lens transform/renderer."""
data: LogitLensData | None = None


def _format_lens(
logits: torch.Tensor,
tokenizer,
Expand Down Expand Up @@ -171,20 +163,22 @@ def _run_causal_mediation(
logits = torch.cat(per_layer_logits, dim=0).save()

if remote and backend is not None:
return {"job_id": backend.job_id}
# Nothing to return: the values land later, on the backend's stream.
return None

return {"logits": logits}


@router.post("/start", response_model=CausalMediationResponse)
async def start_causal_mediation(
@router.post("/run")
async def run_causal_mediation(
req: CausalMediationRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
"""Patch one residual across prompts and lens the result, streamed (see ``sse``)."""
model = state[req.model]
_validate_indices(req, model)
backend = state.make_backend(model=model)
backend = state.make_backend(model)

raw = _run_causal_mediation(
model,
Expand All @@ -198,53 +192,23 @@ async def start_causal_mediation(
backend=backend,
)

if "job_id" in raw:
return {"job_id": raw["job_id"]}

input_tokens = _decode_input_tokens(model.tokenizer, req.tgt_prompt)
data = _format_lens(
raw["logits"],
tokenizer=model.tokenizer,
model_name=req.model,
input_tokens=input_tokens,
n_layers=model.num_layers,
top_k=req.topk,
include_entropy=req.include_entropy,
)
return {"data": data}


@router.post("/results/{job_id}", response_model=CausalMediationResponse)
async def collect_causal_mediation(
job_id: str,
req: CausalMediationRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
backend = state.make_backend(job_id=job_id)
results = backend()

# The model can be deregistered from the catalog (NDIF stopped serving it)
# between /start and /results; state[...] raises KeyError in that case.
# Surface a clear 503 instead of an opaque 500.
try:
model = state[req.model]
except KeyError:
raise HTTPException(
status_code=503,
detail=f"Model {req.model} is no longer available; please re-run.",
)
# Read off the model now rather than when the values land: one connection
# holds the request open, so unlike the old collect step there is no window
# in which NDIF could stop serving this model and leave `state[...]` raising.
tokenizer = model.tokenizer
input_tokens = _decode_input_tokens(tokenizer, req.tgt_prompt)

data = _format_lens(
results["logits"],
tokenizer=tokenizer,
model_name=req.model,
input_tokens=input_tokens,
n_layers=model.num_layers,
top_k=req.topk,
include_entropy=req.include_entropy,
)
def process(saves: dict):
return _format_lens(
saves["logits"],
tokenizer=tokenizer,
model_name=req.model,
input_tokens=input_tokens,
n_layers=model.num_layers,
top_k=req.topk,
include_entropy=req.include_entropy,
)

return {"data": data}
# `_run_causal_mediation` returns None when remote -- the values come off the
# backend's stream -- and the saved values themselves when local.
return stream(backend if state.remote else raw, process)
51 changes: 15 additions & 36 deletions workbench/_api/routes/j_lens.py
Original file line number Diff line number Diff line change
@@ -1,54 +1,33 @@
from fastapi import APIRouter, Depends
from pydantic import BaseModel
from ..state import AppState, get_state
from ..auth import require_user_email

from ..data_models import NDIFResponse

from nnsightful.types import JLensData
from nnsightful.tools.j_lens import j_lens

from ..auth import require_user_email
from ..sse import stream_tool
from ..state import AppState, get_state

router = APIRouter()


class JLensRequest(BaseModel):
model: str
prompt: str
topk: int = 5 # Number of top-k predictions per cell
include_entropy: bool = True # Whether to include entropy data


class JLensResponse(NDIFResponse):
data: JLensData | None = None


@router.post("/start", response_model=JLensResponse)
async def start_j_lens(
@router.post("/run")
async def run_j_lens(
req: JLensRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
model = state[req.model]
backend = state.make_backend(model=model)

output = j_lens._run(model, req.prompt, remote=state.remote, backend=backend, non_blocking=state.remote, raw=False, top_k=req.topk)

if not backend.blocking:
return {"job_id": output}


return {"data": j_lens.to_data_obj(**output)}


@router.post("/results/{job_id}", response_model=JLensResponse)
async def collect_j_lens(
job_id: str,
req: JLensRequest,
state: AppState = Depends(get_state),
user_email: str = Depends(require_user_email),
):
backend = state.make_backend(job_id=job_id)
results = backend()['results']

data = j_lens.to_data_obj(**results)

return {"data": data}
"""Run the Jacobian lens, streaming status until the data lands (see ``sse``)."""
return stream_tool(
state,
j_lens,
state[req.model],
req.prompt,
top_k=req.topk,
)
Loading
Loading