43 lines
869 B
Python
43 lines
869 B
Python
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
from typing import Protocol
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class InferenceTiming:
|
|
preprocess_ms: float
|
|
generate_ms: float
|
|
decode_ms: float
|
|
total_ms: float
|
|
|
|
|
|
class Grounder(Protocol):
|
|
model_path: str
|
|
|
|
def generate(
|
|
self,
|
|
image_path: str,
|
|
prompt: str,
|
|
max_new_tokens: int = 128,
|
|
) -> str: ...
|
|
|
|
def generate_with_metrics(
|
|
self,
|
|
image_path: str,
|
|
prompt: str,
|
|
max_new_tokens: int = 128,
|
|
) -> tuple[str, InferenceTiming, dict[str, float]]: ...
|
|
|
|
def cuda_memory_metrics(self) -> dict[str, float]: ...
|
|
|
|
def implementation_metadata(self) -> dict[str, str]: ...
|
|
|
|
|
|
class TextGenerator(Protocol):
|
|
def generate_text(
|
|
self,
|
|
prompt: str,
|
|
max_new_tokens: int = 256,
|
|
) -> str: ...
|