steering-showcase / gpu_entry.py
wang2226's picture
Fix what a check of the Space found (615cffb)
54d246f verified
Raw History Blame Contribute Delete
3.01 kB
"""The single GPU entry point: measure one scenario at all nine steering strengths.
On ZeroGPU, ``@spaces.GPU`` attaches a GPU only for the duration of this call. The arguments and the
returned trace are pickled, so the trace holds plain Python types. One call measures every value of
alpha, so the slider never needs the GPU again.
"""
from __future__ import annotations
import os
import threading
import time
from dataclasses import replace
import gradio as gr
import spaces
import models
from steering import trace as T
from steering.core import ALPHAS, measure_states, steering_vector
from steering.scenarios import BY_ID, tuning_for
#: Tokens generated per state. The recorded results use the same limit.
MAX_NEW_TOKENS = int(os.environ.get("MAX_NEW_TOKENS", "40"))
#: The steering hook is registered on a shared model for the length of a measurement, so two
#: measurements in one process would steer each other. On ZeroGPU each call runs in its own worker
#: process and this lock is never contended; elsewhere it makes concurrent runs take turns.
_MODEL_LOCK = threading.Lock()
# A run takes 7 to 9 s on the cards this Space has been given. ZeroGPU multiplies the duration by a
# factor for the card (1.5 here) and refuses a visitor whose remaining quota is below the product, so
# 20 s asks for 30 s: over twice the slowest run, and a smaller request also keeps the queue priority up.
@spaces.GPU(duration=20)
def measure(model_key: str, scenario_id: str, prompt: str, prefix: str) -> dict:
"""Runs the model. The only function that touches a GPU, so the only one that costs quota."""
t0 = time.perf_counter()
loaded = models.REGISTRY[model_key]
spec = BY_ID[scenario_id]
model_id = loaded.spec.model_id
layer, scale = tuning_for(model_id, spec, loaded.depth)
custom = prompt != spec.prompt or prefix != spec.prefix
try:
with _MODEL_LOCK:
vector = steering_vector(
loaded.tok, loaded.model, prompt, prefix,
spec.negative_examples, spec.positive_examples, layer,
)
states = measure_states(
loaded.tok, loaded.model, prompt, prefix, vector, layer,
scale=scale, alphas=ALPHAS, generate=True, max_new_tokens=MAX_NEW_TOKENS, top_k=512,
)
except ValueError as exc:
# ZeroGPU passes a gr.Error through with its message and reduces anything else to its class
# name, so say what happened in words a visitor can act on.
raise gr.Error(f"This text could not be measured along the {spec.title} direction: {exc}") from None
return T.build(
replace(spec, prompt=prompt, prefix=prefix),
states,
model_id=model_id,
layer=layer,
scale=scale,
depth=loaded.depth,
token_limit=MAX_NEW_TOKENS,
source="live",
takeaway=T.live_takeaway(spec, custom),
custom=custom,
env=T.environment(loaded.model, gpu_seconds=time.perf_counter() - t0),
)