Skip to content

topicjev.gen

Embeddings and structured text generation. See the Embeddings and LLM Generation guides.

Embeddings

Embedder

Embedder(
    mname: str,
    *,
    batch_size: int = 128,
    prefix: str = "",
    **kwargs: Any,
)

Bases: ModelBackend

Base class for text embedding backends.

Source code in topicjev/gen/embed.py
def __init__(
    self,
    mname: str,
    *,
    batch_size: int = 128,
    prefix: str = "",
    **kwargs: Any,
) -> None:
    super().__init__(mname, **kwargs)
    self.bsize: int = batch_size
    self.prefix: str = prefix
    self.save_dir: Path | None = (
        Path(kwargs.get("save_dir", None))
        if kwargs.get("save_dir", None) is not None
        else None
    )
    self.index: faiss.IndexFlatIP | None = None

embed_docs abstractmethod

embed_docs(docs: list[str]) -> ndarray

Embed a list of documents into a 2D numpy array of shape (n_docs, dim).

Source code in topicjev/gen/embed.py
@abstractmethod
def embed_docs(self, docs: list[str]) -> np.ndarray:
    """Embed a list of documents into a 2D numpy array of shape (n_docs, dim)."""

get_embeds

get_embeds(docs: list[str]) -> ndarray

Compute embeddings for a list of documents.

Source code in topicjev/gen/embed.py
def get_embeds(self, docs: list[str]) -> np.ndarray:
    """Compute embeddings for a list of documents."""
    if self.save_dir is not None and (self.save_dir / "embeds.bin").exists():
        self.index = faiss.read_index(str(self.save_dir / "embeds.bin"))
        return np.array(self.index.reconstruct_n(0, len(docs)))

    for i in tqdm(range(0, len(docs), self.bsize), desc="Embedding docs"):
        batch = docs[i : i + self.bsize]
        embeds = self.embed_docs(batch)
        if self.index is None:
            self.index = faiss.IndexFlatIP(embeds.shape[1])
        self.index.add(embeds)

    if self.save_dir is not None:
        self.save_dir.mkdir(parents=True, exist_ok=True)
        faiss.write_index(self.index, str(self.save_dir / "embeds.bin"))

    return np.array(self.index.reconstruct_n(0, len(docs)))

LocalEmbedder

LocalEmbedder(mname: str, **kwargs: Any)

Bases: Embedder

Embeds text using local SentenceTransformer models.

Source code in topicjev/gen/embed.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.dev, self.dtype = detect_device()
    self.norm_embeds: bool = kwargs.get("normalize_embeddings", True)
    self.trust_remote_code: bool = kwargs.get("trust_remote_code", True)

load_model

load_model() -> None

Load SentenceTransformer model on detected device.

Source code in topicjev/gen/embed.py
def load_model(self) -> None:
    """Load SentenceTransformer model on detected device."""
    from sentence_transformers import SentenceTransformer

    self.model = SentenceTransformer(
        self.mname,
        trust_remote_code=self.trust_remote_code,
        model_kwargs={"torch_dtype": self.dtype},
    ).to(self.dev)
    self.model.eval()

embed_docs

embed_docs(docs: list[str]) -> ndarray

Compute document embeddings.

Source code in topicjev/gen/embed.py
def embed_docs(self, docs: list[str]) -> np.ndarray:
    """Compute document embeddings."""
    if self.model is None:
        self.load_model()
    formatted_docs = [f"{self.prefix}{d}" for d in docs] if self.prefix else docs
    with torch.inference_mode():
        embed_batch = self.model.encode(
            formatted_docs,
            batch_size=8,
            normalize_embeddings=self.norm_embeds,
            show_progress_bar=False,
        )
    return embed_batch

close

close() -> None

Offload embedding model to CPU and clear GPU memory.

Source code in topicjev/gen/embed.py
def close(self) -> None:
    """Offload embedding model to CPU and clear GPU memory."""
    if self.model is not None and hasattr(self.model, "cpu"):
        try:
            self.model.cpu()
        except Exception as e:
            print("Failed to move model to cpu during close: %s", e)
    super().close()
    empty_device_cache()

Generation

Generator

Generator(
    mname: str,
    *,
    temperature: float = 0.1,
    max_tokens: int = 4000,
    max_input_tokens: int = 4000,
    json_mode: bool = True,
    reason: bool = False,
    batch_size: int = 4,
    max_retries: int = 5,
    **kwargs: Any,
)

Bases: ModelBackend

Base class for prompt-to-text generation and structured JSON extraction.

Source code in topicjev/gen/chat.py
def __init__(
    self,
    mname: str,
    *,
    temperature: float = 0.1,
    max_tokens: int = 4000,
    max_input_tokens: int = 4000,
    json_mode: bool = True,
    reason: bool = False,
    batch_size: int = 4,
    max_retries: int = 5,
    **kwargs: Any,
) -> None:
    super().__init__(mname, **kwargs)
    self.temp = temperature
    self.max_tokens = max_tokens
    self.max_inpt_tokens = max_input_tokens
    self.json_mode = json_mode
    self.reason = reason
    self.bsize = batch_size
    self.max_retries = max_retries

query abstractmethod

query(
    user_prompt: str,
    sys_prompt: str = "You are a helpful AI assistant",
) -> Union[str, dict[str, Any]]

Query the model with a single prompt and optional system instructions.

Source code in topicjev/gen/chat.py
@abstractmethod
def query(
    self,
    user_prompt: str,
    sys_prompt: str = "You are a helpful AI assistant",
) -> Union[str, dict[str, Any]]:
    """Query the model with a single prompt and optional system instructions."""

batch abstractmethod

batch(
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]

Execute a prompt template over a batch of input documents.

Source code in topicjev/gen/chat.py
@abstractmethod
def batch(
    self,
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]:
    """Execute a prompt template over a batch of input documents."""

run_chain

run_chain(
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]

Alias for batch() for pipeline compatibility.

Source code in topicjev/gen/chat.py
def run_chain(
    self,
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]:
    """Alias for `batch()` for pipeline compatibility."""
    return self.batch(prompt, input_docs, req_keys=req_keys)

LocalGenerator

LocalGenerator(mname: str, **kwargs: Any)

Bases: CausalLMBackend, Generator

Generates text using a local causal language model on hardware devices.

Source code in topicjev/gen/chat.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.tokenizer: PreTrainedTokenizerBase | Any = None

tok property writable

tok: PreTrainedTokenizerBase | Any

Alias for tokenizer for consistency across local backends.

load_model

load_model() -> None

Load causal LM, left-padded tokenizer, and optionally compile for CUDA.

Source code in topicjev/gen/chat.py
def load_model(self) -> None:
    """Load causal LM, left-padded tokenizer, and optionally compile for CUDA."""
    super().load_model()
    self.tokenizer = self.tok
    self.tokenizer.clean_up_tokenization_spaces = False
    if self.dev.type == "cuda":
        self.model = torch.compile(self.model, mode="default", dynamic=True)

batch_gen

batch_gen(prompts: list[str]) -> list[str]

Generate text outputs for a list of formatted prompts.

Source code in topicjev/gen/chat.py
def batch_gen(self, prompts: list[str]) -> list[str]:
    """Generate text outputs for a list of formatted prompts."""
    res: list[str] = []
    with torch.inference_mode():
        for i in tqdm(range(0, len(prompts), self.bsize), desc="Generating"):
            batch = prompts[i : i + self.bsize]
            enc = self.tokenizer(
                batch,
                padding=True,
                truncation=True,
                max_length=self.max_inpt_tokens,
                return_tensors="pt",
            ).to(self.dev)
            gen_kwargs: dict[str, Any] = {
                "max_new_tokens": self.max_tokens,
                "pad_token_id": self.tokenizer.pad_token_id,
                "eos_token_id": self.tokenizer.eos_token_id,
                "use_cache": True,
            }
            if self.temp > 0.0:
                gen_kwargs["do_sample"] = True
                gen_kwargs["temperature"] = self.temp
                gen_kwargs["top_k"] = 20
                gen_kwargs["top_p"] = 0.95
            else:
                gen_kwargs["do_sample"] = False

            outputs = self.model.generate(**enc, **gen_kwargs)
            gen_tokens = outputs[:, enc["input_ids"].shape[1] :]
            decoded = self.tokenizer.batch_decode(
                gen_tokens, skip_special_tokens=True
            )
            res.extend(decoded)
    return res

query

query(
    user_prompt: str,
    sys_prompt: str = "You are a helpful AI assistant",
) -> Union[str, dict[str, Any]]

Run single-prompt inference.

Source code in topicjev/gen/chat.py
def query(
    self,
    user_prompt: str,
    sys_prompt: str = "You are a helpful AI assistant",
) -> Union[str, dict[str, Any]]:
    """Run single-prompt inference."""
    if self.model is None:
        self.load_model()
    formatted_prompt = self._wrap(user_prompt, sys_prompt)
    raw_result = self.batch_gen([formatted_prompt])[0]
    resp_text = raw_result.strip()
    return clean_json(resp_text) if self.json_mode else str(resp_text)

batch

batch(
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]

Batch-generate responses across input documents.

Source code in topicjev/gen/chat.py
def batch(
    self,
    prompt: str,
    input_docs: list[Any],
    req_keys: list[str] | None = None,
) -> list[Any]:
    """Batch-generate responses across input documents."""
    if self.model is None:
        self.load_model()

    if req_keys is not None and self.json_mode:
        return self._retry_chain(prompt, input_docs, req_keys)

    formatted_prompts = [self._wrap(prompt.format(**doc)) for doc in input_docs]
    raw_results = self.batch_gen(formatted_prompts)

    if self.json_mode:
        return [clean_json(res.strip()) for res in raw_results]
    return [res.strip() for res in raw_results]

close

close() -> None

Offload local model and clear GPU cache.

Source code in topicjev/gen/chat.py
def close(self) -> None:
    """Offload local model and clear GPU cache."""
    if self.model is not None and hasattr(self.model, "cpu"):
        try:
            self.model.cpu()
        except Exception as e:
            print("Failed to move model to cpu during close: %s", e)
    super().close()
    self.tokenizer = None
    empty_device_cache()

clean_json

clean_json(resp: str) -> dict[str, Any]

Extract and parse a JSON dictionary from an LLM response string.

Source code in topicjev/gen/chat.py
def clean_json(resp: str) -> dict[str, Any]:
    """Extract and parse a JSON dictionary from an LLM response string."""
    try:
        resp_text = str(resp).strip()
        if "```" in resp_text:
            match = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", resp_text, re.DOTALL)
            if match:
                resp_text = match.group(1).strip()
            else:
                resp_text = resp_text.replace("```json", "").replace("```", "").strip()
        parsed = json.loads(resp_text)
        return parsed if isinstance(parsed, dict) else {}
    except (json.JSONDecodeError, TypeError, ValueError):
        try:
            start = resp.find("{")
            end = resp.rfind("}")
            if start != -1 and end != -1 and end > start:
                parsed = json.loads(resp[start : end + 1])
                return parsed if isinstance(parsed, dict) else {}
        except Exception:
            pass
        return {}