"""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), )