lightrft.datasets.sft_dataset¶
- class lightrft.datasets.sft_dataset.SFTDataset(*args: Any, **kwargs: Any)[source]¶
Bases:
DatasetDataset for SFT model
- collate_fn(item_list)[source]¶
Collate function to batch samples with padding.
- Parameters:
item_list (list) – List of tuples (prompt_ids_len, input_id, attention_mask, info).
- Returns:
Batched tensors (prompt_ids_lens, input_ids, attention_masks, infos).
- Return type:
tuple
- packing_collate_fn(item_list)[source]¶
Collate function for packing multiple samples into a single sequence.
- Parameters:
item_list (list) – List of tuples (prompt_ids_len, input_id, _, info).
- Returns:
Packed tensors (packed_input_ids, packed_attention_masks, prompt_ids_lens, infos).
- Return type:
tuple
- lightrft.datasets.sft_dataset.preprocess_data(data, input_template=None, input_key='input', output_key=None, apply_chat_template=None, multiturn=False)[source]¶
Preprocess data sample into prompt and response strings.
- Parameters:
data (dict) – Raw data sample dictionary.
input_template (Optional[str]) – Optional template string to format the input.
input_key (str) – Key for input field in data.
output_key (Optional[str]) – Key for output field in data (None for pretrain mode).
apply_chat_template (Optional[Callable]) – Optional chat template function.
multiturn (bool) – Whether to handle multi-turn conversations.
- Returns:
Tuple of (prompt, response) strings.
- Return type:
Tuple[str, str]