diff --git a/plugins/online-data-mixing/README.md b/plugins/online-data-mixing/README.md index 8644440f..657dd524 100644 --- a/plugins/online-data-mixing/README.md +++ b/plugins/online-data-mixing/README.md @@ -71,6 +71,11 @@ Rewards | Description `TRAIN_LOSS` | Training loss where loss is maintained across categories and is updated based on the latest loss and sampled dataset/category. Higher values mean requirement of more samples. `VALIDATION_LOSS` | Validation loss across categories calculated using evaluation datasets from each of the categories. Higher values mean requirement of more samples. `GRADNORM` | Gradient norm where norms are maintained across categories and are updated based on the latest values and sampled dataset/category. Higher values mean reducing samples from that particular dataset/category. +`LEARNABILITY` | Compares model loss on a zero-shot eval batch against a few-shot (templated) batch for the same category: `1 - loss_few_shot / loss_zero_shot`. Higher values mean the category benefits more from in-context examples, i.e. there is more to learn from it. Requires `templated_eval_dataset_dict` (see below). +`VELOCITY` | Tracks how quickly a category's eval loss is dropping between consecutive reward computations: `1 - current_loss / previous_loss`. Higher values mean the category is still improving quickly and should keep being sampled. +`COMBINED` | Exponentially-decayed blend of `LEARNABILITY` and `VELOCITY`: `R(t) = alpha(t) * learnability + (1 - alpha(t)) * velocity`, `alpha(t) = exp(-beta * t / total_steps)`. Favors learnability early in training and velocity later. Requires `templated_eval_dataset_dict` (see below). + +`LEARNABILITY` and `COMBINED` require passing `templated_eval_dataset_dict` (and `templated_eval_collators_dict`) to `OnlineMixingDataset` — a few-shot-templated version of `eval_dataset_dict`, keyed by the same category names, used as the "few-shot" side of the learnability comparison. Constructing `OnlineMixingDataset` with one of these reward types but without `templated_eval_dataset_dict` raises a `ValueError`. ### Adding a Custom Reward diff --git a/plugins/online-data-mixing/src/fms_acceleration_odm/odm/dataset.py b/plugins/online-data-mixing/src/fms_acceleration_odm/odm/dataset.py index 6257c964..655fa3bc 100644 --- a/plugins/online-data-mixing/src/fms_acceleration_odm/odm/dataset.py +++ b/plugins/online-data-mixing/src/fms_acceleration_odm/odm/dataset.py @@ -37,6 +37,9 @@ def __init__( reward_type=Reward.ENTROPY, auto_categorize_config: Optional[dict | AutoCategorizeConfig] = None, seed: Optional[int] = 42, + templated_eval_dataset_dict: Optional[DatasetDict] = None, + templated_eval_collators_dict: Optional[dict] = None, + beta: float = 1.0, ): """Mixes datasets with sampling ratios learnt using Multi Armed Bandit (MAB) EXP3 and rewards defined. @@ -72,7 +75,30 @@ def __init__( has only one key. seed (Optional[int], optional): Base seed for the dataset-level RNG so all distributed ranks iterate over the exact same sample order. Defaults to 42. + templated_eval_dataset_dict (Optional[DatasetDict], optional): keys are category + names and values are HF eval datasets rendered with few-shot templates. Required + (together with `templated_eval_collators_dict`) when `reward_type` is + Reward.LEARNABILITY or Reward.COMBINED, used as the "few-shot" side of the + learnability comparison against `eval_dataset_dict` (the "zero-shot" side). + templated_eval_collators_dict (Optional[dict], optional): collator corresponding + to each dataset in `templated_eval_dataset_dict`. + beta (float, optional): decay hyper-parameter for Reward.COMBINED's alpha(t) + schedule. Defaults to 1.0. """ + # should be one of Reward + self.reward_type = reward_type + if isinstance(self.reward_type, str): + self.reward_type = self.reward_type.upper() + self.reward_type = Reward[self.reward_type] + + if self.reward_type in (Reward.LEARNABILITY, Reward.COMBINED) and ( + templated_eval_dataset_dict is None or templated_eval_collators_dict is None + ): + raise ValueError( + "templated_eval_dataset_dict and templated_eval_collators_dict must both " + "be provided when reward_type is Reward.LEARNABILITY or Reward.COMBINED." + ) + self.auto_categorize = len(dataset_dict.keys()) == 1 self._auto_categorize_config = self._build_auto_categorize_config( auto_categorize_config @@ -83,6 +109,14 @@ def __init__( eval_dataset_dict, eval_collators_dict = self._maybe_auto_categorize_dataset( eval_dataset_dict, eval_collators_dict, dataset_role="eval" ) + if templated_eval_dataset_dict is not None: + templated_eval_dataset_dict, templated_eval_collators_dict = ( + self._maybe_auto_categorize_dataset( + templated_eval_dataset_dict, + templated_eval_collators_dict, + dataset_role="templated_eval", + ) + ) logger.info( """Values set to OnlineMixingDataset @@ -122,6 +156,10 @@ def __init__( self.eval_collators_dict = eval_collators_dict self.eval_dataset_dict = eval_dataset_dict self.eval_dataset_dict_dl = {} + self.templated_eval_collators_dict = templated_eval_collators_dict + self.templated_eval_dataset_dict = templated_eval_dataset_dict + self.templated_eval_dataset_dict_dl = {} + self.beta = beta # iterators of the dataloaders self.train_dataset_dict_dl_iter = {} # to reset iterators to dataloaders @@ -356,6 +394,27 @@ def _reset_eval_dataloaders(self): else None ) + self.templated_eval_dataset_dict_dl = {} + if self.templated_eval_dataset_dict: + for k, _ in self.templated_eval_dataset_dict.items(): + self.templated_eval_dataset_dict_dl[k] = ( + iter( + DataLoader( + self.templated_eval_dataset_dict[k], + self.eval_batch_size, + shuffle=True, + num_workers=0, + collate_fn=( + self.templated_eval_collators_dict[k] + if self.templated_eval_collators_dict + else None + ), + ) + ) + if self.templated_eval_dataset_dict[k] + else None + ) + def _build_auto_categorize_config(self, config): if isinstance(config, AutoCategorizeConfig): return config @@ -454,13 +513,10 @@ def _extract_information_from_state_for_reward(self, state=None, category=None): Returns: dict: arguments prepared for compute_reward function """ - if state is None: - return {} - if self.reward_type.startswith(Reward.ENTROPY): + if state is None or self.reward_type.startswith(Reward.ENTROPY): return {} - if self.reward_type == Reward.TRAIN_LOSS: - return {"train_loss_history": [d for d in state.log_history if "loss" in d]} - if self.reward_type == Reward.VALIDATION_LOSS: + + def _validation_loss_info(): assert category is not None return { "eval_loss_history": [ @@ -469,11 +525,22 @@ def _extract_information_from_state_for_reward(self, state=None, category=None): if f"eval_{category}_loss" in d ] } - if self.reward_type == Reward.GRADNORM: - return { + + extractors = { + Reward.TRAIN_LOSS: lambda: { + "train_loss_history": [d for d in state.log_history if "loss" in d] + }, + Reward.VALIDATION_LOSS: _validation_loss_info, + Reward.GRADNORM: lambda: { "gradnorm_history": [d for d in state.log_history if "grad_norm" in d] - } - return {} + }, + Reward.COMBINED: lambda: { + "train_step": state.global_step, + "total_steps": getattr(state, "max_steps", None), + "beta": self.beta, + }, + } + return extractors.get(self.reward_type, dict)() def update_sampling_weights(self, model, accelerator, state): """Function to update MAB weights based on the reward type provided @@ -493,6 +560,8 @@ def update_sampling_weights(self, model, accelerator, state): rewards = [0] * self.total_categories count = [0] * self.total_categories eval_dataset_dict = {} + templated_eval_dataset_dict = {} + needs_templated = self.reward_type in (Reward.LEARNABILITY, Reward.COMBINED) device = accelerator.device if accelerator else torch.device(0) self._reset_eval_dataloaders() for c in range(self.total_categories): @@ -503,10 +572,22 @@ def update_sampling_weights(self, model, accelerator, state): if self.eval_dataset_dict_dl.get(self.id2cat[c], None) else None ) + if needs_templated: + templated_eval_dataset_dict[self.id2cat[c]] = ( + accelerator.prepare( + self.templated_eval_dataset_dict_dl[self.id2cat[c]] + ) + if self.templated_eval_dataset_dict_dl.get(self.id2cat[c], None) + else None + ) else: eval_dataset_dict[self.id2cat[c]] = self.eval_dataset_dict_dl.get( self.id2cat[c], None ) + if needs_templated: + templated_eval_dataset_dict[self.id2cat[c]] = ( + self.templated_eval_dataset_dict_dl.get(self.id2cat[c], None) + ) for c in tqdm( range(self.total_categories), total=self.total_categories, desc="Categories" ): # for trian loss you dont need to iterate over eval dataset. @@ -525,6 +606,34 @@ def update_sampling_weights(self, model, accelerator, state): ) rewards[c] += rc count[c] += 1 + elif needs_templated: + for batch, templated_batch in tqdm( + zip( + eval_dataset_dict[self.id2cat[c]], + templated_eval_dataset_dict[self.id2cat[c]], + ), + desc="Reward computation over eval dataset", + ): + zero_shot_batch = {k: v.to(device) for k, v in batch.items()} + few_shot_batch = { + k: v.to(device) for k, v in templated_batch.items() + } + rc = compute_reward( + model=model, + batch=zero_shot_batch, + vocab_size=32000, + reward_type=self.reward_type, + current_category=c, + total_categories=self.total_categories, + last_sampled_category=self.arm_idx, + zero_shot_batch=zero_shot_batch, + few_shot_batch=few_shot_batch, + **self._extract_information_from_state_for_reward( + state, self.id2cat[c] + ), + ) + rewards[c] += rc + count[c] += zero_shot_batch["input_ids"].shape[0] else: for batch in tqdm( eval_dataset_dict[self.id2cat[c]], diff --git a/plugins/online-data-mixing/src/fms_acceleration_odm/odm/reward.py b/plugins/online-data-mixing/src/fms_acceleration_odm/odm/reward.py index 3404dd5d..67f89a64 100644 --- a/plugins/online-data-mixing/src/fms_acceleration_odm/odm/reward.py +++ b/plugins/online-data-mixing/src/fms_acceleration_odm/odm/reward.py @@ -2,13 +2,17 @@ # Standard from enum import StrEnum, auto +from logging import getLogger from typing import Dict +import math # Third Party from transformers import PreTrainedModel import torch import torch.nn.functional as F +logger = getLogger(__name__) + class Reward(StrEnum): ENTROPY = auto() @@ -17,6 +21,9 @@ class Reward(StrEnum): TRAIN_LOSS = auto() VALIDATION_LOSS = auto() GRADNORM = auto() + LEARNABILITY = auto() + VELOCITY = auto() + COMBINED = auto() # TODO: @@ -31,6 +38,8 @@ class Reward(StrEnum): GRADNORM_DATA = {"buffer": []} +VELOCITY_DATA = {"buffer": []} + def compute_reward( model: PreTrainedModel, @@ -43,6 +52,11 @@ def compute_reward( last_sampled_category=None, total_categories=None, current_category=None, + zero_shot_batch=None, + few_shot_batch=None, + train_step=None, + total_steps=None, + beta: float = 1.0, ) -> float: """ Compute rewards based on the provided reward_type. @@ -67,6 +81,25 @@ def compute_reward( Similar to TRAIN_LOSS reward here we use overall gradnorm. However, Higher grad norm categories should be less priortized. + Learnability reward: LEARNABILITY + Compares the model's loss on a zero-shot batch against a few-shot (templated) batch + for the same category: reward = 1 - loss_few_shot / loss_zero_shot. Higher values mean + the category benefits more from in-context examples, i.e. the model has more to learn + from that category. + + Velocity reward: VELOCITY + Tracks how quickly a category's (eval) loss is dropping between consecutive reward + computations: reward = 1 - current_loss / previous_loss. Higher values mean the category + is still improving quickly and should keep being sampled. + + Combined reward: COMBINED + Exponentially-decayed blend of LEARNABILITY and VELOCITY: + R(t) = alpha(t) * learnability + (1 - alpha(t)) * velocity + alpha(t) = exp(-beta * train_step / total_steps) + Early in training the blend favors LEARNABILITY (which category teaches the model the + most relative to its current baseline); later it favors VELOCITY (which category is + still yielding gains in absolute loss). + Args: model (PreTrainedModel): HF Model object batch (torch.Tensor): Batch of samples (input_ids, labels, attention_mask) @@ -78,71 +111,172 @@ def compute_reward( last_sampled_category: index of the last sampled category total_categories: total number of categories current_category: currently being reward computed category + zero_shot_batch: batch of zero-shot samples, used by LEARNABILITY/COMBINED + few_shot_batch: batch of few-shot (templated) samples, used by LEARNABILITY/COMBINED + train_step: current training step, used by COMBINED to compute alpha(t) + total_steps: total number of training steps, used by COMBINED to compute alpha(t) + beta: decay hyper-parameter for COMBINED's alpha(t). Defaults to 1.0. Returns: float """ if reward_type.startswith(Reward.ENTROPY): - with torch.inference_mode(): - outputs = model(**batch) - shift_logits = outputs.logits[:, :-1, :] + return _compute_entropy_reward(model, batch, vocab_size, reward_type) + + handlers = { + Reward.TRAIN_LOSS: lambda: _compute_train_loss_reward( + train_loss_history, last_sampled_category, current_category, total_categories + ), + Reward.VALIDATION_LOSS: lambda: _compute_validation_loss_reward( + eval_loss_history, current_category, total_categories + ), + Reward.GRADNORM: lambda: _compute_gradnorm_reward( + gradnorm_history, last_sampled_category, current_category, total_categories + ), + Reward.LEARNABILITY: lambda: _compute_learnability_reward( + model, zero_shot_batch, few_shot_batch + ), + Reward.VELOCITY: lambda: _compute_velocity_reward( + model, batch, current_category, total_categories + ), + Reward.COMBINED: lambda: _compute_combined_reward( + model, + batch, + current_category, + total_categories, + zero_shot_batch, + few_shot_batch, + train_step, + total_steps, + beta, + ), + } + if reward_type not in handlers: + raise TypeError(f"Reward {reward_type} not supported") + return handlers[reward_type]() + + +def _compute_entropy_reward(model, batch, vocab_size, reward_type) -> float: + with torch.inference_mode(): + outputs = model(**batch) + shift_logits = outputs.logits[:, :-1, :] + + log_probs = F.log_softmax(shift_logits, dim=-1) + probs = torch.exp(log_probs) + + entropy = -torch.sum(probs * log_probs, dim=-1) + sum_p_log_sq = torch.sum(probs * (log_probs**2), dim=-1) + varentropy = sum_p_log_sq - (entropy**2) + + entropy_last_token = entropy[:, -1] + + mask = batch["attention_mask"][:, 1:] + + entropy = (entropy * mask).sum(dim=-1) / mask.sum(dim=-1) + varentropy = (varentropy * mask).sum(dim=-1) / mask.sum(dim=-1) - log_probs = F.log_softmax(shift_logits, dim=-1) - probs = torch.exp(log_probs) + max_entropy = torch.log( + torch.tensor(vocab_size, dtype=entropy.dtype, device=entropy.device) + ) - entropy = -torch.sum(probs * log_probs, dim=-1) - sum_p_log_sq = torch.sum(probs * (log_probs**2), dim=-1) - varentropy = sum_p_log_sq - (entropy**2) + entropy = (entropy / max_entropy).clamp(0.0, 1.0) + varentropy = (varentropy / max_entropy**2).clamp(0.0, 1.0) + entropy_last_token = (entropy_last_token / max_entropy).clamp(0.0, 1.0) + if reward_type == Reward.ENTROPY: + return entropy.sum().item() + if reward_type == Reward.ENTROPY3_VARENT1: + return 0.75 * entropy.sum().item() + 0.25 * varentropy.sum().item() + return entropy_last_token.sum().item() - entropy_last_token = entropy[:, -1] - mask = batch["attention_mask"][:, 1:] +def _compute_train_loss_reward( + train_loss_history, last_sampled_category, current_category, total_categories +) -> float: + if not train_loss_history: + raise ValueError("train_loss_history cannot be a empty list or None") + if not TRAIN_LOSS_DATA["buffer"]: + TRAIN_LOSS_DATA["buffer"] = [1e-100] * total_categories + TRAIN_LOSS_DATA["buffer"][last_sampled_category] = train_loss_history[-1]["loss"] + return TRAIN_LOSS_DATA["buffer"][current_category] - entropy = (entropy * mask).sum(dim=-1) / mask.sum(dim=-1) - varentropy = (varentropy * mask).sum(dim=-1) / mask.sum(dim=-1) - max_entropy = torch.log( - torch.tensor(vocab_size, dtype=entropy.dtype, device=entropy.device) +def _compute_validation_loss_reward( + eval_loss_history, current_category, total_categories +) -> float: + if not eval_loss_history: + raise ValueError( + "eval_loss_history cannot be a empty list or None." + "Make sure you are using eval_strategy and eval_steps" + "allowing atleast 1 evaluation before reward computation." ) + if not EVAL_LOSS_DATA["buffer"]: + EVAL_LOSS_DATA["buffer"] = [1e-100] * total_categories + EVAL_LOSS_DATA["buffer"][current_category] = eval_loss_history[-1]["loss"] + return EVAL_LOSS_DATA["buffer"][current_category] - entropy = (entropy / max_entropy).clamp(0.0, 1.0) - varentropy = (varentropy / max_entropy**2).clamp(0.0, 1.0) - entropy_last_token = (entropy_last_token / max_entropy).clamp(0.0, 1.0) - if reward_type == Reward.ENTROPY: - return entropy.sum().item() - if reward_type == Reward.ENTROPY3_VARENT1: - return 0.75 * entropy.sum().item() + 0.25 * varentropy.sum().item() - if reward_type == Reward.ENTROPY_LAST_TOKEN: - return entropy_last_token.sum().item() - if reward_type == Reward.TRAIN_LOSS: - if not train_loss_history: - raise ValueError("train_loss_history cannot be a empty list or None") - if not TRAIN_LOSS_DATA["buffer"]: - TRAIN_LOSS_DATA["buffer"] = [1e-100] * total_categories - TRAIN_LOSS_DATA["buffer"][last_sampled_category] = train_loss_history[-1][ - "loss" - ] - return TRAIN_LOSS_DATA["buffer"][current_category] - if reward_type == Reward.VALIDATION_LOSS: - if not eval_loss_history: - raise ValueError( - "eval_loss_history cannot be a empty list or None." - "Make sure you are using eval_strategy and eval_steps" - "allowing atleast 1 evaluation before reward computation." - ) - if not EVAL_LOSS_DATA["buffer"]: - EVAL_LOSS_DATA["buffer"] = [1e-100] * total_categories - EVAL_LOSS_DATA["buffer"][current_category] = eval_loss_history[-1]["loss"] - return EVAL_LOSS_DATA["buffer"][current_category] - if reward_type == Reward.GRADNORM: - if not gradnorm_history: - raise ValueError( - "gradnorm_history cannot be a empty list or None." - "Make sure grad norm is made available." - ) - if not GRADNORM_DATA["buffer"]: - GRADNORM_DATA["buffer"] = [1e-100] * total_categories - GRADNORM_DATA["buffer"][last_sampled_category] = 1 / ( - gradnorm_history[-1]["grad_norm"] + 0.0001 + +def _compute_gradnorm_reward( + gradnorm_history, last_sampled_category, current_category, total_categories +) -> float: + if not gradnorm_history: + raise ValueError( + "gradnorm_history cannot be a empty list or None." + "Make sure grad norm is made available." ) - return GRADNORM_DATA["buffer"][current_category] - raise TypeError(f"Reward {reward_type} not supported") + if not GRADNORM_DATA["buffer"]: + GRADNORM_DATA["buffer"] = [1e-100] * total_categories + GRADNORM_DATA["buffer"][last_sampled_category] = 1 / ( + gradnorm_history[-1]["grad_norm"] + 0.0001 + ) + return GRADNORM_DATA["buffer"][current_category] + + +def _compute_combined_reward( + model, + batch, + current_category, + total_categories, + zero_shot_batch, + few_shot_batch, + train_step, + total_steps, + beta, +) -> float: + learn_r = _compute_learnability_reward(model, zero_shot_batch, few_shot_batch) + vel_r = _compute_velocity_reward(model, batch, current_category, total_categories) + if not total_steps: + logger.warning( + "COMBINED reward received an empty total_steps; falling back to " + "alpha=1.0 (pure LEARNABILITY) for this call." + ) + alpha = 1.0 + else: + alpha = math.exp(-beta * train_step / total_steps) + return alpha * learn_r + (1.0 - alpha) * vel_r + + +def _compute_learnability_reward(model, zero_shot_batch, few_shot_batch) -> float: + if zero_shot_batch is None or few_shot_batch is None: + raise ValueError( + "zero_shot_batch and few_shot_batch cannot be None for LEARNABILITY/" + "COMBINED rewards." + ) + with torch.inference_mode(): + loss_zero_shot = model(**zero_shot_batch).loss.item() + loss_few_shot = model(**few_shot_batch).loss.item() + if loss_zero_shot == 0: + return 0.0 + return 1.0 - loss_few_shot / loss_zero_shot + + +def _compute_velocity_reward(model, batch, current_category, total_categories) -> float: + if batch is None: + raise ValueError("batch cannot be None for VELOCITY/COMBINED rewards.") + with torch.inference_mode(): + current_loss = model(**batch).loss.item() + if not VELOCITY_DATA["buffer"]: + VELOCITY_DATA["buffer"] = [None] * total_categories + previous_loss = VELOCITY_DATA["buffer"][current_category] + VELOCITY_DATA["buffer"][current_category] = current_loss + if not previous_loss: + return 0.0 + return 1.0 - current_loss / previous_loss diff --git a/plugins/online-data-mixing/tests/test_compute_reward.py b/plugins/online-data-mixing/tests/test_compute_reward.py index 39d5dd2a..659bdf54 100644 --- a/plugins/online-data-mixing/tests/test_compute_reward.py +++ b/plugins/online-data-mixing/tests/test_compute_reward.py @@ -181,3 +181,112 @@ def test_compute_reward( gradnorm_history=h, ) assert returned_reward == r, f"expected {r} but got {returned_reward}" + + +def _make_batch(batch_size, seq_length, vocab_size, offset=0): + input_ids = ( + torch.arange(offset, offset + batch_size * seq_length).reshape( + batch_size, seq_length + ) + % vocab_size + ) + attention_mask = torch.ones(batch_size, seq_length, dtype=torch.long) + return {"input_ids": input_ids, "labels": input_ids, "attention_mask": attention_mask} + + +def test_compute_reward_learnability(): + loaded_model = AutoModelForCausalLM.from_pretrained("Maykeye/TinyLLama-v0") + zero_shot_batch = _make_batch(3, 6, 50) + few_shot_batch = _make_batch(3, 10, 50, offset=1) + + reward = compute_reward( + model=loaded_model, + batch=None, + vocab_size=50, + reward_type=Reward.LEARNABILITY, + current_category=0, + total_categories=2, + zero_shot_batch=zero_shot_batch, + few_shot_batch=few_shot_batch, + ) + + with torch.inference_mode(): + expected = 1.0 - ( + loaded_model(**few_shot_batch).loss.item() + / loaded_model(**zero_shot_batch).loss.item() + ) + assert reward == pytest.approx(expected) + + with pytest.raises(ValueError): + compute_reward( + model=loaded_model, + batch=None, + vocab_size=50, + reward_type=Reward.LEARNABILITY, + current_category=0, + total_categories=2, + ) + + +def test_compute_reward_velocity(): + loaded_model = AutoModelForCausalLM.from_pretrained("Maykeye/TinyLLama-v0") + batch = _make_batch(3, 6, 50) + + first_reward = compute_reward( + model=loaded_model, + batch=batch, + vocab_size=50, + reward_type=Reward.VELOCITY, + current_category=0, + total_categories=2, + ) + # no prior loss recorded for this category yet + assert first_reward == 0.0 + + second_reward = compute_reward( + model=loaded_model, + batch=batch, + vocab_size=50, + reward_type=Reward.VELOCITY, + current_category=0, + total_categories=2, + ) + # same batch twice through an unchanging model -> loss unchanged -> no velocity + assert second_reward == pytest.approx(0.0) + + with pytest.raises(ValueError): + compute_reward( + model=loaded_model, + batch=None, + vocab_size=50, + reward_type=Reward.VELOCITY, + current_category=1, + total_categories=2, + ) + + +def test_compute_reward_combined(): + loaded_model = AutoModelForCausalLM.from_pretrained("Maykeye/TinyLLama-v0") + zero_shot_batch = _make_batch(3, 6, 50) + few_shot_batch = _make_batch(3, 10, 50, offset=1) + + with torch.inference_mode(): + loss_zero_shot = loaded_model(**zero_shot_batch).loss.item() + loss_few_shot = loaded_model(**few_shot_batch).loss.item() + expected_learnability = 1.0 - loss_few_shot / loss_zero_shot + + # total_steps falsy -> falls back to alpha=1.0 (pure learnability) + reward = compute_reward( + model=loaded_model, + batch=zero_shot_batch, + vocab_size=50, + reward_type=Reward.COMBINED, + current_category=0, + total_categories=2, + zero_shot_batch=zero_shot_batch, + few_shot_batch=few_shot_batch, + train_step=0, + total_steps=None, + beta=1.0, + ) + assert reward == pytest.approx(expected_learnability) diff --git a/plugins/online-data-mixing/tests/test_online_data.py b/plugins/online-data-mixing/tests/test_online_data.py index 1f3ba187..7a13235a 100644 --- a/plugins/online-data-mixing/tests/test_online_data.py +++ b/plugins/online-data-mixing/tests/test_online_data.py @@ -13,7 +13,9 @@ # limitations under the License. # Third Party +from datasets import Dataset from torch.utils.data import IterableDataset +from transformers import AutoModelForCausalLM # pylint: disable=import-error import pytest @@ -98,3 +100,86 @@ def test_online_data_mix_learning( assert sum(x == y for x, y in zip(categories_chosen, expected_arm_idx)) >= ( len(expected_arm_idx) / 2 ), "Not even half of the choices were correct" + + +def _hf_dataset(seq_length, vocab_size, num_rows=4): + return Dataset.from_dict( + { + "input_ids": [[i % vocab_size] * seq_length for i in range(num_rows)], + "attention_mask": [[1] * seq_length for _ in range(num_rows)], + "labels": [[i % vocab_size] * seq_length for i in range(num_rows)], + } + ).with_format("torch") + + +def test_online_data_learnability_requires_templated_eval_dataset(): + seq_length = 6 + vocab_size = 50 + train_data_dict = { + "data_1": _hf_dataset(seq_length, vocab_size), + "data_2": _hf_dataset(seq_length, vocab_size), + } + collators_dict = {"data_1": None, "data_2": None} + with pytest.raises(ValueError): + OnlineMixingDataset( + train_data_dict, + collators_dict, + train_data_dict, + collators_dict, + output_dir="odm", + reward_type=Reward.LEARNABILITY, + ) + + +@pytest.mark.parametrize("reward_type", [Reward.LEARNABILITY, Reward.COMBINED]) +def test_online_data_update_sampling_weights_with_templated_eval_dataset(reward_type): + seq_length = 6 + vocab_size = 50 + train_data_dict = { + "data_1": _hf_dataset(seq_length, vocab_size), + "data_2": _hf_dataset(seq_length, vocab_size), + } + collators_dict = {"data_1": None, "data_2": None} + templated_data_dict = { + "data_1": _hf_dataset(seq_length + 4, vocab_size), + "data_2": _hf_dataset(seq_length + 4, vocab_size), + } + + dataset = OnlineMixingDataset( + train_data_dict, + collators_dict, + train_data_dict, + collators_dict, + output_dir="odm", + reward_type=reward_type, + eval_batch_size=2, + templated_eval_dataset_dict=templated_data_dict, + templated_eval_collators_dict=collators_dict, + beta=1.0, + ) + + class DummyState: + global_step = 0 + max_steps = 10 + log_history = [] + + class CPUAccelerator: + device = torch.device("cpu") + + def prepare(self, x): + return x + + def reduce(self, x, reduction): # pylint: disable=unused-argument + return x + + model = AutoModelForCausalLM.from_pretrained("Maykeye/TinyLLama-v0") + # update_sampling_weights() moves eval batches to accelerator.device (or + # torch.device(0), i.e. cuda:0, if no accelerator is given). Pass a + # single-process CPU stub so this test doesn't require a GPU. + dataset.update_sampling_weights(model, accelerator=CPUAccelerator(), state=DummyState()) + + assert dataset.log["rewards"], "expected rewards to be logged after update" + counts = list(dataset.log["count"]) + assert all( + count > 0 for count in counts + ), "expected every category to accumulate a nonzero count"