Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
7e767b2
Add forward for seq2seq tokenization
alex-jw-brooks Sep 18, 2023
4c899dc
Add comparator test for seq2seq forwarding
alex-jw-brooks Sep 24, 2023
a159209
Split seq2seq tokenization preprocessing
alex-jw-brooks Sep 24, 2023
ad323a9
Add forward to seq2seq tokenization (no batch)
alex-jw-brooks Sep 24, 2023
0d894ac
Add batch forwarding tests for seq2seq/causal lm
alex-jw-brooks Sep 24, 2023
fe2c57c
Add batch forward for causal lm / seq2seq
alex-jw-brooks Sep 24, 2023
727e784
rewrite causal lm tok tests to check chunking
alex-jw-brooks Sep 25, 2023
ac9fedc
Implement chunked tokenization for causal lm
alex-jw-brooks Sep 25, 2023
5fc9849
linting, formatting
alex-jw-brooks Sep 25, 2023
c514d68
Turn on seq2seq tokenization by default
alex-jw-brooks Sep 26, 2023
0b75f80
Hack: explictly force stream unwrapping to false
alex-jw-brooks Sep 26, 2023
52a33a5
Approximate port of old causal lm tokenization logic
alex-jw-brooks Sep 27, 2023
9f1aff3
Hack - use default collator for causal LM
alex-jw-brooks Sep 28, 2023
33e3e14
Add generic test for left/right padding causal lm seq approach
alex-jw-brooks Sep 28, 2023
a515e35
Do left / right padding via tokenizer pad
alex-jw-brooks Sep 28, 2023
73e804e
Add simple tests for default data collator
alex-jw-brooks Sep 29, 2023
349cbdc
Update concat seq test for corrected padding
alex-jw-brooks Sep 29, 2023
ed1f8e6
Fix legacy ported sequence length bug
alex-jw-brooks Sep 29, 2023
579db79
Update comments for tokenizer changes
alex-jw-brooks Sep 29, 2023
560d5bc
linting and formatting
alex-jw-brooks Sep 29, 2023
511e561
Update causal lm docstrings and type hints
alex-jw-brooks Sep 29, 2023
ccd9c13
Fix remainder handling in chunking test
alex-jw-brooks Sep 29, 2023
52f9910
Add chunk example, use extend
alex-jw-brooks Oct 2, 2023
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
22 changes: 5 additions & 17 deletions caikit_nlp/modules/text_generation/peft_prompt_tuning.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,7 @@
from torch.optim import AdamW
from torch.utils.data import DataLoader
from tqdm import tqdm
from transformers import (
AutoModelForCausalLM,
DataCollatorForLanguageModeling,
default_data_collator,
)
from transformers import AutoModelForCausalLM, default_data_collator
from transformers.models.auto.tokenization_auto import AutoTokenizer
from transformers.optimization import get_linear_schedule_with_warmup
import numpy as np
Expand Down Expand Up @@ -890,12 +886,8 @@ def _get_collate_fn(tokenizer: AutoTokenizer, task_type: str) -> Callable:
Callable
collate_fn to be used for processing batches from our datasets.
"""
if task_type == "CAUSAL_LM":
return DataCollatorForLanguageModeling(
tokenizer=tokenizer,
return_tensors="pt",
mlm=False,
)
# HACK: Do NOT use the causal LM collator (for now) because
# want to set labels ourselves. TODO: centralize collator management.
return default_data_collator

@staticmethod
Expand Down Expand Up @@ -936,15 +928,11 @@ def _get_data_loaders_from_stream(
torch.utils.data.DataLoader
DataLoader to be used for training / evaluating the stream data.
"""
(
tokenize_function,
requires_unwrapping,
) = base_model.build_task_tokenize_closure(
(tokenize_function, _,) = base_model.build_task_tokenize_closure(
tokenizer, max_source_length, max_target_length, verbalizer, task_ids=0
)
mapped_stream = train_stream.map(tokenize_function)
if requires_unwrapping:
mapped_stream = mapped_stream.flatten()
# TODO: Deprecate and remove stream wrapper & use trainer
wrapped_stream = SimpleIterableStreamWrapper(mapped_stream, shuffle=shuffle)
dataloader = DataLoader(
wrapped_stream, collate_fn=collate_fn, batch_size=batch_size
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -593,7 +593,7 @@ def _preprocess_function(
mapped_dataset = dataset.map(
base_model.tokenize_function,
fn_kwargs=fn_kwargs,
batched=base_model.REQUIRES_TOKEN_UNWRAPPING,
batched=False,
Comment thread
alex-jw-brooks marked this conversation as resolved.
# Drop the input / output columns; we need to do this for dimensions to play
# happily when operating on batched inputs for causal language modeling.
remove_columns=["input", "output"],
Expand Down
Loading