pandora_llm.features.LossRatio

Module Contents

class pandora_llm.features.LossRatio.LossRatio(model_name, ref_model_name, model_revision=None, model_cache_dir=None, ref_model_revision=None, ref_model_cache_dir=None)[source]

Bases: pandora_llm.features.base.FeatureComputer, pandora_llm.features.base.LLMHandler

Computes likelihood ratio against a reference model (also known as a reference-based attack). Mathematically, this is log-likelihood from primary model minus log-likelihood from reference model

Parameters:
  • model_name (str)

  • ref_model_name (str)

  • model_revision (str)

  • model_cache_dir (str)

  • ref_model_revision (str)

  • ref_model_cache_dir (str)

ref_model_name[source]
ref_model_revision = None[source]
ref_model_cache_dir = None[source]
load_model(stage)[source]

Loads model into memory

Parameters:

stage (str) – ‘primary’ or ‘ref’

Return type:

None

compute_features(dataloader, accelerator)[source]

Computes the negative log-likelihood (NLL) feature for the given dataloader.

Parameters:
  • dataloader (torch.utils.data.DataLoader) – The dataloader providing input sequences.

  • accelerator (accelerate.Accelerator) – The Accelerator object for distributed or mixed-precision training.

Returns:

The NLL feature for each sequence in the dataloader.

Raises:

Exception – If the model is not loaded before calling this method.

Return type:

jaxtyping.Float[torch.Tensor, n]

static reduce(primary_log_probs, ref_log_probs)[source]

Computes loss ratio by computing primary_log_probs-ref_log_probs

Parameters:
  • primary_log_probs (jaxtyping.Float[torch.Tensor, n]) – Log probs from primary model

  • ref_log_probs (jaxtyping.Float[torch.Tensor, n]) – Log probs from reference model

Returns:

primary-ref log probs

Return type:

jaxtyping.Float[torch.Tensor, n]