Skip to content

TextGradientTrainingDataset

Bases: GradientTrainingDataset

Text modality wrapper for GradientTrainingDataset.

Supplies padding from tokenizer (pad_token_id for 'input_ids') and cache_key_fields ['input_text', 'label'] when caching is used. Both keys must be present in the batch.

Source code in gradiend/trainer/text/common/dataset.py
def __init__(
    self,
    training_data: Any,
    tokenizer: Any,
    gradient_creator: Any,
    *,
    source: str = 'factual',
    target: str = 'diff',
    cache_dir: Optional[str] = None,
    use_cached_gradients: bool = True,
    dtype: torch.dtype = torch.float32,
    device: Optional[torch.device] = None,
    return_metadata: bool = False,
    timing_steps: int = 0,
    timing_label: str = "text-gradient",
):
    pad_token_id = getattr(tokenizer, 'pad_token_id', 0) if tokenizer is not None else 0

    def get_padding_value(subkey: str) -> int:
        return pad_token_id if 'input_ids' in subkey else 0

    super().__init__(
        training_data,
        gradient_creator,
        source=source,
        target=target,
        cache_dir=cache_dir,
        use_cached_gradients=use_cached_gradients,
        cache_key_fields=self.CACHE_KEY_FIELDS if (cache_dir and use_cached_gradients) else None,
        dtype=dtype,
        device=device,
        return_metadata=return_metadata,
        get_padding_value=get_padding_value,
        timing_steps=timing_steps,
        timing_label=timing_label,
    )
    self.tokenizer = tokenizer

CACHE_KEY_FIELDS class-attribute instance-attribute

CACHE_KEY_FIELDS = ['input_text', 'label']

tokenizer instance-attribute

tokenizer = tokenizer