Skip to content

topicjev.entail

Zero-shot classification backends. The Zero-shot Classification guide explains how to choose a backend and lists the options that each one reads from its keyword arguments.

Base classes

Entailment

Entailment(
    mname: str | None,
    *,
    batch_size: int = 16,
    max_tokens: int = 512,
    multi_lbl: bool = False,
    decoys: list[str] | None = None,
    probes: list[str] | None = None,
    temperature: float = 1.0,
    threshold: float | None = None,
    **kwargs,
)

Bases: ModelBackend

Base class for entailment and zero-shot classification backends.

Orchestrates pairing, probe calibration, temperature scaling, normalization, and threshold gating across heterogeneous model backends.

Source code in topicjev/entail/base.py
def __init__(
    self,
    mname: str | None,
    *,
    batch_size: int = 16,
    max_tokens: int = 512,
    multi_lbl: bool = False,
    decoys: list[str] | None = None,
    probes: list[str] | None = None,
    temperature: float = 1.0,
    threshold: float | None = None,
    **kwargs,
) -> None:
    super().__init__(mname, **kwargs)
    self.bsize = batch_size
    self.max_tokens = max_tokens
    self.multi_lbl = multi_lbl
    self.decoys: list[str] = list(decoys) if decoys else []
    self.probes: list[str] = list(probes) if probes else []
    self.temp = temperature
    self.threshold = threshold
    self.str_name = "base"
    self.jev = False

    self.model: PreTrainedModel | Any = None
    self.tok: PreTrainedTokenizerBase | Any = None

prep_pairs abstractmethod

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Format input documents and labels into scoring pairs.

Source code in topicjev/entail/base.py
@abstractmethod
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Format input documents and labels into scoring pairs."""

batch_score abstractmethod

batch_score(
    pairs: list[Pair],
) -> list[float] | list[list[float]]

Score input pairs and return raw, uncalibrated scores.

Source code in topicjev/entail/base.py
@abstractmethod
def batch_score(self, pairs: list[Pair]) -> list[float] | list[list[float]]:
    """Score input pairs and return raw, uncalibrated scores."""

entail

entail(
    docs: list[str], lbls: list[str]
) -> list[dict[str, Any]]

Score documents against labels and return classification results.

Source code in topicjev/entail/base.py
def entail(self, docs: list[str], lbls: list[str]) -> list[dict[str, Any]]:
    """Score documents against labels and return classification results."""
    prompts = self.probes + docs
    all_lbls = list(lbls) + list(self.decoys)
    pairs = self.prep_pairs(prompts, all_lbls)

    raw_scores = self.batch_score(pairs)
    scores = torch.as_tensor(raw_scores, dtype=torch.float32).reshape(
        len(prompts), len(all_lbls)
    )

    if self.probes:
        n = len(self.probes)
        scores = scores[n:] - scores[:n].mean(dim=0, keepdim=True)

    scores = scores / self.temp

    if not self.multi_lbl or self.jev:
        scores = torch.softmax(scores, dim=1)
    else:
        scores = torch.sigmoid(scores)

    return EntailResults(
        lbls=lbls,
        multi_lbl=self.multi_lbl,
        decoys=self.decoys,
        threshold=self.threshold,
    ).compute_results(scores.tolist())

close

close() -> None

Release allocated model and tokenizer resources.

Source code in topicjev/entail/base.py
def close(self) -> None:
    """Release allocated model and tokenizer resources."""
    super().close()
    self.tok = None

LocalEntailment

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

Bases: Entailment

Entailment backend that loads local PyTorch models onto hardware devices.

Source code in topicjev/entail/base.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.dev, self.dtype = detect_device()

max_len property

max_len: int

Maximum context length supported by tokenizer or model configuration.

close

close() -> None

Offload local model and clear GPU cache.

Source code in topicjev/entail/base.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()
    empty_device_cache()

RemoteEntailment

RemoteEntailment(
    mname: str | None,
    *,
    batch_size: int = 16,
    max_tokens: int = 512,
    multi_lbl: bool = False,
    decoys: list[str] | None = None,
    probes: list[str] | None = None,
    temperature: float = 1.0,
    threshold: float | None = None,
    **kwargs,
)

Bases: Entailment

Entailment backend for hosted remote API services without local GPU requirements.

Source code in topicjev/entail/base.py
def __init__(
    self,
    mname: str | None,
    *,
    batch_size: int = 16,
    max_tokens: int = 512,
    multi_lbl: bool = False,
    decoys: list[str] | None = None,
    probes: list[str] | None = None,
    temperature: float = 1.0,
    threshold: float | None = None,
    **kwargs,
) -> None:
    super().__init__(mname, **kwargs)
    self.bsize = batch_size
    self.max_tokens = max_tokens
    self.multi_lbl = multi_lbl
    self.decoys: list[str] = list(decoys) if decoys else []
    self.probes: list[str] = list(probes) if probes else []
    self.temp = temperature
    self.threshold = threshold
    self.str_name = "base"
    self.jev = False

    self.model: PreTrainedModel | Any = None
    self.tok: PreTrainedTokenizerBase | Any = None

Local backends

XEncoderEntail

XEncoderEntail(mname: str, **kwargs)

Bases: LocalEntailment

Zero-shot classification using Natural Language Inference (NLI) cross-encoders.

Pairs documents as premises and hypothesis templates ('The text discusses {label}').

Source code in topicjev/entail/xencoder.py
def __init__(self, mname: str, **kwargs) -> None:
    super().__init__(mname, **kwargs)
    self._idx: tuple[int, int, int | None] | None = None
    self.hyp: str = kwargs.get("hyp", TOPIC_HYP)
    self.str_name = "xencoder"

entail_indices property

entail_indices: tuple[int, int, int | None]

Resolve indices for entailment, contradiction, and neutral output labels.

load_model

load_model() -> None

Load cross-encoder classification model and tokenizer.

Source code in topicjev/entail/xencoder.py
def load_model(self) -> None:
    """Load cross-encoder classification model and tokenizer."""
    from transformers import AutoModelForSequenceClassification, AutoTokenizer

    self.model = AutoModelForSequenceClassification.from_pretrained(
        self.mname,
        dtype=self.dtype,
    ).to(self.dev)
    self.tok = AutoTokenizer.from_pretrained(self.mname)
    self.model.eval()

prep_pairs

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Format documents and hypothesis labels into cross-encoder premise-hypothesis pairs.

Source code in topicjev/entail/xencoder.py
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Format documents and hypothesis labels into cross-encoder premise-hypothesis pairs."""
    longest_hyp = max(
        (len(self.tok(self.hyp.format(l)).input_ids) for l in lbls), default=0
    )
    room = max(1, self.max_len - longest_hyp)
    ids = self.tok([str(p) for p in prompts], truncation=True, max_length=room)
    docs = self.tok.batch_decode(ids.input_ids, skip_special_tokens=True)
    return [Pair(d, self.hyp.format(l)) for d in docs for l in lbls]

batch_score

batch_score(pairs: list[Pair]) -> list[float]

Compute cross-encoder classification logits.

Source code in topicjev/entail/xencoder.py
def batch_score(self, pairs: list[Pair]) -> list[float]:
    """Compute cross-encoder classification logits."""
    scores: list[float] = []
    ent, cnt, neu = self.entail_indices
    space_fn = self._space_fn(neu)
    with torch.no_grad():
        for i in tqdm(range(0, len(pairs), self.bsize), desc="Entailment scoring"):
            batch: list[Pair] = pairs[i : i + self.bsize]
            enc = self.tok(
                [p.inpt for p in batch],
                [p.targ for p in batch],
                padding="longest",
                truncation="only_first",
                max_length=self.max_len,
                return_tensors="pt",
            )
            enc = {k: v.to(self.dev) for k, v in enc.items()}
            lgts = self.model(**enc).logits.float()
            out = space_fn(lgts, ent, cnt, neu)
            scores.extend(out.cpu().tolist())
    return scores

Seq2SeqEntail

Seq2SeqEntail(mname: str, **kwargs)

Bases: LocalEntailment

Scores candidate labels by their average per-token log-likelihood under teacher-forcing.

Source code in topicjev/entail/seq2seq.py
def __init__(self, mname: str, **kwargs) -> None:
    super().__init__(mname, **kwargs)
    self.templ: str = kwargs.get("template", CLASS_CUE)
    self.res_tokens: int = kwargs.get("reserve", 50)
    self.str_name = "seq2seq"

load_model

load_model() -> None

Load encoder-decoder seq2seq model and tokenizer.

Source code in topicjev/entail/seq2seq.py
def load_model(self) -> None:
    """Load encoder-decoder seq2seq model and tokenizer."""
    from transformers import AutoModelForSeq2SeqLM, AutoTokenizer

    self.model = AutoModelForSeq2SeqLM.from_pretrained(
        self.mname,
        dtype=self.dtype,
        attn_implementation="sdpa",
    ).to(self.dev)
    self.tok = AutoTokenizer.from_pretrained(self.mname)
    self.model.eval()

prep_pairs

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Format input prompts and candidate target labels into teacher-forced pairs.

Source code in topicjev/entail/seq2seq.py
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Format input prompts and candidate target labels into teacher-forced pairs."""
    overhead = len(self.tok(self.templ.format("")).input_ids)
    room = max(1, self.max_len - overhead - self.res_tokens)
    ids = self.tok([str(d) for d in prompts], truncation=True, max_length=room)
    docs = self.tok.batch_decode(ids.input_ids, skip_special_tokens=True)
    return [Pair(self.templ.format(d), targ=l) for d in docs for l in lbls]

batch_score

batch_score(pairs: list[Pair]) -> list[float]

Compute mean target token log-probabilities under teacher forcing.

Source code in topicjev/entail/seq2seq.py
def batch_score(self, pairs: list[Pair]) -> list[float]:
    """Compute mean target token log-probabilities under teacher forcing."""
    scores: list[float] = []
    with torch.no_grad():
        for i in tqdm(range(0, len(pairs), self.bsize), desc="Entailment scoring"):
            batch: list[Pair] = pairs[i : i + self.bsize]
            enc = self.tok(
                [p.inpt for p in batch],
                padding=True,
                truncation=True,
                max_length=self.max_len,
                return_tensors="pt",
            )
            tgt = self.tok(
                [p.targ for p in batch], padding=True, return_tensors="pt"
            )
            enc = {k: v.to(self.dev) for k, v in enc.items()}
            ids = tgt["input_ids"].to(self.dev)
            mask = tgt["attention_mask"].to(self.dev)
            dec_in = self.model._shift_right(ids)
            lgts = self.model(**enc, decoder_input_ids=dec_in).logits.float()
            logprobs = torch.log_softmax(lgts, dim=-1)
            top_lk = logprobs.gather(-1, ids.unsqueeze(-1)).squeeze(-1) * mask
            scores.extend(
                (top_lk.sum(-1) / mask.sum(-1).clamp(min=1)).cpu().tolist()
            )
    return scores

GoalExEntail

GoalExEntail(mname: str, **kwargs)

Bases: YesNoLogitMixin, Seq2SeqEntail

Evaluates candidate labels as independent binary property verification questions.

Scores the logit difference between 'Yes' and 'No' tokens at the initial decoder step of an encoder-decoder seq2seq model in a single forward pass without autoregressive generation.

Inspired by the GoalEx approach (Wang, Shang, and Zhong, 2023)_.

.. _(Wang, Shang, and Zhong, 2023): https://arxiv.org/abs/2305.13749

Source code in topicjev/entail/goalex.py
def __init__(self, mname: str, **kwargs) -> None:
    super().__init__(mname, **kwargs)
    self.templ: str = kwargs.get("template", GOALEX)
    self.temp: float = kwargs.get("temperature", 0.1)
    self.res_tokens: int = kwargs.get("reserve", 8)
    self.str_name = "goalex"

prep_pairs

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Format input prompts and labels into GoalEx question pairs.

Source code in topicjev/entail/goalex.py
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Format input prompts and labels into GoalEx question pairs."""
    blank = len(self.tok(self.templ.format(text="", property="")).input_ids)
    longest = max((len(self.tok(str(l)).input_ids) for l in lbls), default=0)
    room = max(1, self.max_len - blank - longest - self.res_tokens)
    ids = self.tok([str(d) for d in prompts], truncation=True, max_length=room)
    docs = self.tok.batch_decode(ids.input_ids, skip_special_tokens=True)
    return [Pair(self.templ.format(text=d, property=l)) for d in docs for l in lbls]

batch_score

batch_score(pairs: list[Pair]) -> list[float]

Score pairs by extracting yes/no logits at initial decoder token position.

Source code in topicjev/entail/goalex.py
def batch_score(self, pairs: list[Pair]) -> list[float]:
    """Score pairs by extracting yes/no logits at initial decoder token position."""
    scores: list[float] = []
    start_id = getattr(self.model.config, "decoder_start_token_id", None)
    if start_id is None:
        start_id = self.tok.pad_token_id
    with torch.no_grad():
        for i in tqdm(range(0, len(pairs), self.bsize), desc="Entailment scoring"):
            batch: list[Pair] = pairs[i : i + self.bsize]
            enc = self.tok(
                [p.inpt for p in batch],
                padding="longest",
                truncation=True,
                max_length=self.max_len,
                return_tensors="pt",
            )
            enc = {k: v.to(self.dev) for k, v in enc.items()}
            dec = torch.full(
                (enc["input_ids"].size(0), 1),
                start_id,
                dtype=torch.long,
                device=self.dev,
            )
            lgts = self.model(**enc, decoder_input_ids=dec).logits.float()[:, 0, :]
            scores.extend(self._yes_no_score(lgts).cpu().tolist())
    return scores

CausalEntail

CausalEntail(mname: str, **kwargs)

Bases: CausalLMBackend, GoalExEntail

Evaluates binary property verification from the last token of a causal LM forward pass.

Requires left-padding so position index -1 corresponds to the terminal prompt token. Inherits model loading and chat templating from CausalLMBackend.

Source code in topicjev/entail/causal.py
def __init__(self, mname: str, **kwargs) -> None:
    super().__init__(mname, **kwargs)
    self.chat: bool = kwargs.get("chat", True)
    self.thinking: bool = kwargs.get("thinking", False)
    self.min_len: int = kwargs.get("min_len", 0)
    self.num_beams: int = kwargs.get("num_beams", 1)
    self.str_name = "causal"

prep_pairs

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Format input prompts and labels into left-padded scoring pairs.

Source code in topicjev/entail/causal.py
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Format input prompts and labels into left-padded scoring pairs."""
    blank = len(
        self.tok(self._wrap(self.templ.format(text="", property=""))).input_ids
    )
    longest = max((len(self.tok(str(l)).input_ids) for l in lbls), default=0)
    room = max(1, self.max_len - blank - longest - self.res_tokens)
    ids = self.tok([str(d) for d in prompts], truncation=True, max_length=room)
    docs = self.tok.batch_decode(ids.input_ids, skip_special_tokens=True)
    return [
        Pair(self._wrap(self.templ.format(text=d, property=l)))
        for d in docs
        for l in lbls
    ]

batch_score

batch_score(pairs: list[Pair]) -> list[float]

Compute yes/no logit difference at terminal token position.

Source code in topicjev/entail/causal.py
def batch_score(self, pairs: list[Pair]) -> list[float]:
    """Compute yes/no logit difference at terminal token position."""
    scores: list[float] = []
    with torch.no_grad():
        for i in tqdm(range(0, len(pairs), self.bsize), desc="Entailment scoring"):
            batch: list[Pair] = pairs[i : i + self.bsize]
            enc = self.tok(
                [p.inpt for p in batch],
                padding="longest",
                truncation=True,
                max_length=self.max_len,
                return_tensors="pt",
            )
            enc = {k: v.to(self.dev) for k, v in enc.items()}
            lgts = self.model(**enc).logits.float()[:, -1, :]
            scores.extend(self._yes_no_score(lgts).cpu().tolist())
    return scores

Laya and TypeSafe

LayaEntail

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

Bases: JevEntail

Zero-shot classification via local Laya agent models.

Source code in topicjev/entail/jevlike.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(**kwargs)
    self.mname = mname
    self.str_name = "laya"
    self.dev, _ = detect_device()

load_model

load_model() -> None

Load local Laya model on detected hardware device.

Source code in topicjev/entail/jevlike.py
def load_model(self) -> None:
    """Load local Laya model on detected hardware device."""
    from laya.agent import load as laya_load

    self.model = laya_load(model_id_or_path=self.mname, device=str(self.dev))

batch_score

batch_score(pairs: list[Pair]) -> list[list[float]]

Score candidate choices in local batches, returning log-probabilities.

Source code in topicjev/entail/jevlike.py
def batch_score(self, pairs: list[Pair]) -> list[list[float]]:
    """Score candidate choices in local batches, returning log-probabilities."""
    scores: list[list[float]] = []
    for i in tqdm(range(0, len(pairs), self.bsize), desc="Laya scoring"):
        batch: list[Pair] = pairs[i : i + self.bsize]
        reqs = [{"text": b.inpt} for b in batch]
        out = self.model.predict_batch(reqs, self.questions)
        probs = [list(o["answers"]["topic"]["probabilities"].values()) for o in out]
        scores.extend(torch.log(torch.tensor(probs).clamp_min(1e-12)).tolist())
    return scores

close

close() -> None

Release local Laya agent resources and clear GPU memory.

Source code in topicjev/entail/jevlike.py
def close(self) -> None:
    """Release local Laya agent resources and clear GPU memory."""
    super().close()
    empty_device_cache()

JevEntail

JevEntail(**kwargs: Any)

Bases: RemoteEntailment

Zero-shot classification via hosted TypeSafe System-One decision endpoints.

Source code in topicjev/entail/jevlike.py
def __init__(self, **kwargs: Any) -> None:
    super().__init__(mname=None, **kwargs)
    self.str_name = "jev"
    self.questions: dict[str, Any] = {}
    self.decoys = [
        *self.decoys,
        json.dumps({"Other": "None of the other categories fit this text"}),
    ]
    self.jev = True
    self.multi_lbl = False

load_model

load_model() -> None

Initialize remote TypeSafe API client.

Source code in topicjev/entail/jevlike.py
def load_model(self) -> None:
    """Initialize remote TypeSafe API client."""
    from typesafe_sdk import TypeSafeClient

    load_dotenv()
    self.model = TypeSafeClient(api_key=os.environ.get("TYPESAFE_API_KEY"))

prep_pairs

prep_pairs(
    prompts: list[str], lbls: list[str]
) -> list[Pair]

Construct question criteria schema and input prompt pairs.

Source code in topicjev/entail/jevlike.py
def prep_pairs(self, prompts: list[str], lbls: list[str]) -> list[Pair]:
    """Construct question criteria schema and input prompt pairs."""
    criteria: dict[str, str] = {}
    for label in lbls:
        try:
            parsed = json.loads(label)
        except json.JSONDecodeError:
            criteria[label] = label
        else:
            criteria.update(parsed)

    self.questions = {
        "topic": {
            "type": "choice",
            "instructions": JEV_PROMPT,
            "criteria": criteria,
        }
    }
    return [Pair(inpt=p) for p in prompts]

batch_score

batch_score(pairs: list[Pair]) -> list[list[float]]

Score candidate choices via remote endpoint calls, returning log-probabilities.

Source code in topicjev/entail/jevlike.py
def batch_score(self, pairs: list[Pair]) -> list[list[float]]:
    """Score candidate choices via remote endpoint calls, returning log-probabilities."""
    scores: list[list[float]] = []

    def map_probs_to_output_map(probs: dict[str, float]) -> list[float]:
        return [
            probs.get(label, 0.0) for label in self.questions["topic"]["criteria"]
        ]

    with self.model as client:
        for i in tqdm(range(len(pairs)), desc="Jev scoring"):
            resp = client.system_one(
                state={"text": pairs[i].inpt}, questions=self.questions
            )
            probs = map_probs_to_output_map(resp.answers["topic"].probabilities)
            scores.append(torch.log(torch.tensor(probs).clamp_min(1e-12)).tolist())
    return scores

Data classes

Pair dataclass

Pair(inpt: str, targ: str | None = None)

A single (input, optional target) pair sent to a scoring model.

EntailRes dataclass

EntailRes(
    class_id: int,
    prob: float,
    probs: list[float],
    other: str | None = None,
)

Entailment classification result for a single document.

to_dict

to_dict() -> dict[str, Any]

Convert result to dictionary representation.

Source code in topicjev/entail/base.py
def to_dict(self) -> dict[str, Any]:
    """Convert result to dictionary representation."""
    return {
        "class": self.class_id,
        "prob": self.prob,
        "other": self.other,
        "probs": self.probs,
    }

EntailResults dataclass

EntailResults(
    lbls: list[str],
    multi_lbl: bool = False,
    decoys: list[str] = list(),
    threshold: float | None = None,
)

Aggregates batch probabilities into top-1 classification predictions.

Applies threshold gating to assign a document to either its top predicted label or falls back to 'Other'.

compute_results

compute_results(
    batch_probs: list[list[float]],
) -> list[dict[str, Any]]

Map per-label probability distributions to classification result dictionaries.

Source code in topicjev/entail/base.py
def compute_results(self, batch_probs: list[list[float]]) -> list[dict[str, Any]]:
    """Map per-label probability distributions to classification result dictionaries."""
    cand = list(self.lbls) + list(self.decoys)
    results: list[dict[str, Any]] = []
    for probs in batch_probs:
        order = sorted(range(len(probs)), key=lambda k: probs[k], reverse=True)
        idx = order[0]
        max_prob = probs[idx]
        in_taxonomy = idx < len(self.lbls)
        class_id = (
            idx if (max_prob > self._threshold and in_taxonomy) else self._otherid
        )
        other = None if in_taxonomy else cand[idx]
        results.append(
            EntailRes(
                class_id=class_id,
                prob=max_prob,
                probs=probs[: len(self.lbls)],
                other=other,
            ).to_dict()
        )
    return results