Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/llmcompressor/transformers/data/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,5 +9,6 @@
from .gsm8k import GSM8KDataset
from .open_platypus import OpenPlatypusDataset
from .peoples_speech import PeoplesSpeech
from .perfectblend import PerfectBlendDataset
from .ultrachat_200k import UltraChatDataset
from .wikitext import WikiTextDataset
70 changes: 70 additions & 0 deletions src/llmcompressor/transformers/data/perfectblend.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
from copy import deepcopy
from typing import TYPE_CHECKING

from loguru import logger

from llmcompressor.transformers.data import TextGenerationDataset
from llmcompressor.typing import Processor

if TYPE_CHECKING:
from llmcompressor.args import DatasetArguments

ROLE_MAP = {"human": "user", "gpt": "assistant"}


@TextGenerationDataset.register(name="perfectblend")
class PerfectBlendDataset(TextGenerationDataset):
"""
Child text generation class for the open-perfectblend dataset

:param dataset_args: configuration settings for dataset loading
:param split: split from dataset to load, for instance `test` or `train[:5%]`
:param processor: processor or tokenizer to use on dataset
"""

DEFAULT_CHAT_TEMPLATE = (
"{% for message in messages %}\n"
"{% if message['role'] == 'user' %}\n"
"{{ '<|user|>\n' + message['content'] + eos_token }}\n"
"{% elif message['role'] == 'system' %}\n"
"{{ '<|system|>\n' + message['content'] + eos_token }}\n"
"{% elif message['role'] == 'assistant' %}\n"
"{{ '<|assistant|>\n' + message['content'] + eos_token }}\n"
"{% endif %}\n"
"{% if loop.last and add_generation_prompt %}\n"
"{{ '<|assistant|>' }}\n{% endif %}\n{% endfor %}"
)

def __init__(
self, dataset_args: "DatasetArguments", split: str, processor: Processor
):
dataset_args = deepcopy(dataset_args)
dataset_args.dataset = "mlabonne/open-perfectblend"
dataset_args.text_column = "conversations"

super().__init__(dataset_args=dataset_args, split=split, processor=processor)

if (
self.tokenizer is not None
and getattr(self.tokenizer, "chat_template", None) is None
):
self.tokenizer.chat_template = self.DEFAULT_CHAT_TEMPLATE
logger.warning(
"tokenizer.chat_template is not set, using default chat template for "
f"{self.__class__.__name__}"
)

def dataset_template(self, sample):
messages = [
{
"role": ROLE_MAP.get(msg["from"], msg["from"]),
"content": msg["value"],
}
for msg in sample["conversations"]
]
Comment thread
kylesayrs marked this conversation as resolved.

return {
"text": self.processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=False
)
}
2 changes: 1 addition & 1 deletion src/llmcompressor/transformers/data/ultrachat_200k.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from llmcompressor.args import DatasetArguments


@TextGenerationDataset.register(name="ultrachat_200k")
@TextGenerationDataset.register(name="ultrachat_200k", alias="ultrachat")
Comment thread
kylesayrs marked this conversation as resolved.
class UltraChatDataset(TextGenerationDataset):
"""
Child text generation class for the Ultra Chat 200k dataset
Expand Down
Loading