Skip to content

topicjev.compress

Document compression backends. The Compression guide compares them and lists the options that each one reads from its keyword arguments.

Compressor

Compressor(
    mname: str,
    *,
    ratio: float = 0.3,
    batch_size: int = 8,
    **kwargs: Any,
)

Bases: ModelBackend

Base class for document summarization and token-level prompt compression.

Provides device management, maximum context length resolution, and resource lifecycle handling across compression strategies.

Source code in topicjev/compress/base.py
def __init__(
    self,
    mname: str,
    *,
    ratio: float = 0.3,
    batch_size: int = 8,
    **kwargs: Any,
) -> None:
    super().__init__(mname, **kwargs)
    self.ratio = ratio
    self.bsize = batch_size
    self.source_len: int = kwargs.get("source_len", 512)
    self.tok: PreTrainedTokenizerBase | Any = None
    self.dev, self.dtype = detect_device()

max_enc_len property

max_enc_len: int

Maximum context length supported by tokenizer or model configuration.

compress abstractmethod

compress(docs: list[str]) -> list[str]

Compress or summarize a list of documents.

Source code in topicjev/compress/base.py
@abstractmethod
def compress(self, docs: list[str]) -> list[str]:
    """Compress or summarize a list of documents."""

close

close() -> None

Offload model to CPU, release tokenizer, and clear GPU cache.

Source code in topicjev/compress/base.py
def close(self) -> None:
    """Offload model to CPU, release tokenizer, 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.tok = None
    empty_device_cache()

CausalCompressor

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

Bases: CausalLMBackend, Compressor

Summarizes text using a causal language model with dynamic input-relative token budgets.

Formats prompts using chat templates (via CausalLMBackend) and restricts generated token count to ratio * prompt_length.

Source code in topicjev/compress/causal.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.source_len: int = kwargs.get("source_len", 2048)
    self.instruct: str = kwargs.get("instruct", INSTRUCT_CLM)
    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)

compress

compress(docs: list[str]) -> list[str]

Summarize a list of documents in batches.

Source code in topicjev/compress/causal.py
def compress(self, docs: list[str]) -> list[str]:
    """Summarize a list of documents in batches."""
    compr_docs: list[str] = []
    i_docs = [self.instruct.format(text=doc) for doc in docs]
    wrapped = [self._wrap(doc) for doc in i_docs]

    with torch.inference_mode():
        for i in tqdm(range(0, len(wrapped), self.bsize), desc="Compressing docs"):
            batch = wrapped[i : i + self.bsize]
            enc = self.tok(
                batch,
                return_tensors="pt",
                padding=True,
                truncation=True,
                max_length=self.max_enc_len,
                add_special_tokens=not self.chat,
            ).to(self.dev)

            prompt_length = enc["input_ids"].shape[1]
            dynamic_max_len = int(prompt_length * self.ratio)
            safe_max_len = max(dynamic_max_len, self.min_len)

            outputs = self.model.generate(
                **enc,
                max_new_tokens=safe_max_len,
                min_new_tokens=self.min_len,
                num_beams=self.num_beams,
                pad_token_id=self.tok.pad_token_id,
                use_cache=True,
            ).cpu()

            gen_tokens = outputs[:, prompt_length:]
            compr_batch = self.tok.batch_decode(
                gen_tokens, skip_special_tokens=True
            )
            compr_docs.extend(compr_batch)
    return compr_docs

CNNCompressor

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

Bases: Compressor

Summarizes text using seq2seq models fine-tuned for summarization (e.g. BART, PEGASUS).

Target length is computed as ratio * max_enc_len.

Source code in topicjev/compress/seq2seq.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.source_len: int = kwargs.get("source_len", 512)
    self.min_len: int = kwargs.get("min_len", 30)
    self.num_beams: int = kwargs.get("num_beams", 4)

max_len property

max_len: int

Target maximum summary generation token length.

load_model

load_model() -> None

Load encoder-decoder seq2seq model and tokenizer.

Source code in topicjev/compress/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()

compress

compress(docs: list[str]) -> list[str]

Summarize documents using beam search generation.

Source code in topicjev/compress/seq2seq.py
def compress(self, docs: list[str]) -> list[str]:
    """Summarize documents using beam search generation."""
    compr_docs: list[str] = []
    with torch.inference_mode():
        for i in tqdm(range(0, len(docs), self.bsize), desc="Compressing docs"):
            batch = docs[i : i + self.bsize]
            enc = self.tok(
                batch,
                return_tensors="pt",
                padding=True,
                truncation=True,
                max_length=self.max_enc_len,
            ).to(self.dev)

            outputs = self.model.generate(
                **enc,
                max_length=self.max_len,
                min_length=self.min_len,
                num_beams=self.num_beams,
                use_cache=True,
            ).cpu()

            compr_batch = self.tok.batch_decode(outputs, skip_special_tokens=True)
            compr_docs.extend(compr_batch)
    return compr_docs

Seq2SeqCompressor

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

Bases: CNNCompressor

Instruction-guided seq2seq summarizer for general instruction-tuned checkpoints (e.g. FLAN-T5).

Source code in topicjev/compress/seq2seq.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.source_len: int = kwargs.get("source_len", 1024)
    self.instruct: str = kwargs.get("instruct", INSTRUCT_S2S)

compress

compress(docs: list[str]) -> list[str]

Format input documents with instruction prefix before running seq2seq generation.

Source code in topicjev/compress/seq2seq.py
def compress(self, docs: list[str]) -> list[str]:
    """Format input documents with instruction prefix before running seq2seq generation."""
    instruct_docs = [self.instruct.format(doc) for doc in docs]
    return super().compress(instruct_docs)

LinguaCompressor

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

Bases: Compressor

Performs token-level prompt pruning using LLMLingua2 without autoregressive text generation (Pan et al., 2024)_.

.. _(Pan et al., 2024): https://arxiv.org/abs/2403.12968

Source code in topicjev/compress/lingua.py
def __init__(self, mname: str, **kwargs: Any) -> None:
    super().__init__(mname, **kwargs)
    self.target_token: int = kwargs.get("target_token", -1)
    self.force_tokens: list[str] = kwargs.get("force_tokens", DEFAULT_FORCE_TOKENS)
    self.force_digits: bool = kwargs.get("force_digits", False)
    self.drop_consecutive: bool = kwargs.get("drop_consecutive", True)

load_model

load_model() -> None

Initialize LLMLingua PromptCompressor.

Source code in topicjev/compress/lingua.py
def load_model(self) -> None:
    """Initialize LLMLingua PromptCompressor."""
    from llmlingua import PromptCompressor

    self.model = PromptCompressor(
        model_name=self.mname,
        use_llmlingua2=True,
        device_map=str(self.dev) if self.dev.type == "cuda" else "cpu",
    )

compress

compress(docs: list[str]) -> list[str]

Prune input documents to the configured compression rate.

Source code in topicjev/compress/lingua.py
def compress(self, docs: list[str]) -> list[str]:
    """Prune input documents to the configured compression rate."""
    compr_docs: list[str] = []
    for doc in tqdm(docs, desc="Compressing docs"):
        res = self.model.compress_prompt(
            str(doc),
            rate=self.ratio,
            target_token=self.target_token,
            force_tokens=self.force_tokens,
            force_reserve_digit=self.force_digits,
            drop_consecutive=self.drop_consecutive,
        )
        compr_docs.append(res["compressed_prompt"])
    return compr_docs