From 1931379d13e0d043fc811bc49f30b48f0b7960e8 Mon Sep 17 00:00:00 2001 From: Amal Joe Date: Mon, 31 Aug 2026 14:45:34 +0530 Subject: [PATCH 1/3] feat: add LEARNABILITY, VELOCITY, and COMBINED rewards to ODM plugin Adds three new reward types to the online-data-mixing plugin's compute_reward dispatch, ported from a learnability/velocity-driven data mixing research prototype: - LEARNABILITY: 1 - loss_few_shot/loss_zero_shot, comparing a category's zero-shot vs few-shot (templated) eval loss. - VELOCITY: 1 - current_loss/previous_loss, tracking how quickly a category's eval loss is still dropping (module-level buffer, no dev set needed beyond the existing eval_dataset_dict). - COMBINED: exponentially-decayed blend of the two, alpha(t) = exp(-beta*t/total_steps), favoring learnability early in training and velocity later. OnlineMixingDataset gains optional templated_eval_dataset_dict / templated_eval_collators_dict / beta constructor params (required together with LEARNABILITY/COMBINED, validated at construction time) and threads them through _reset_eval_dataloaders, update_sampling_weights, and _extract_information_from_state_for_reward. Co-Authored-By: Claude Sonnet 5 Signed-off-by: Amal Joe --- plugins/online-data-mixing/README.md | 5 + .../src/fms_acceleration_odm/odm/dataset.py | 111 ++++++++++++++++++ .../src/fms_acceleration_odm/odm/reward.py | 82 +++++++++++++ .../tests/test_compute_reward.py | 109 +++++++++++++++++ .../tests/test_online_data.py | 75 ++++++++++++ 5 files changed, 382 insertions(+) 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..66f2b1d4 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 @@ -473,6 +532,16 @@ def _extract_information_from_state_for_reward(self, state=None, category=None): return { "gradnorm_history": [d for d in state.log_history if "grad_norm" in d] } + if self.reward_type == Reward.LEARNABILITY: + return {} + if self.reward_type == Reward.VELOCITY: + return {} + if self.reward_type == Reward.COMBINED: + return { + "train_step": state.global_step, + "total_steps": getattr(state, "max_steps", None), + "beta": self.beta, + } return {} def update_sampling_weights(self, model, accelerator, state): @@ -493,6 +562,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 +574,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 +608,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..f533e50d 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,6 +111,11 @@ 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 """ @@ -145,4 +183,48 @@ def compute_reward( gradnorm_history[-1]["grad_norm"] + 0.0001 ) return GRADNORM_DATA["buffer"][current_category] + if reward_type == Reward.LEARNABILITY: + return _compute_learnability_reward(model, zero_shot_batch, few_shot_batch) + if reward_type == Reward.VELOCITY: + return _compute_velocity_reward(model, batch, current_category, total_categories) + if reward_type == Reward.COMBINED: + 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 raise TypeError(f"Reward {reward_type} not supported") + + +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..e76d259b 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,76 @@ 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 = [] + + model = AutoModelForCausalLM.from_pretrained("Maykeye/TinyLLama-v0") + # update_sampling_weights() moves eval batches to torch.device(0) when no + # accelerator is given; match the model to whatever that resolves to here. + model = model.to(torch.device(0)) + dataset.update_sampling_weights(model, accelerator=None, state=DummyState()) + + assert dataset.log["rewards"], "expected rewards to be logged after update" + assert all( + count > 0 for count in dataset.log["count"] + ), "expected every category to accumulate a nonzero count" From 539ee10beed541744cc675f77bcb52c769bb0796 Mon Sep 17 00:00:00 2001 From: Amal Joe Date: Tue, 8 Sep 2026 21:14:01 +0530 Subject: [PATCH 2/3] fix: resolve pylint errors flagged on ODM reward/dataset changes Splits compute_reward() and _extract_information_from_state_for_reward() into per-reward-type helper functions plus a dispatch table, bringing each back under pylint's too-many-return-statements threshold (both had grown to 9 returns after the LEARNABILITY/VELOCITY/COMBINED additions). Also fixes a not-an-iterable false positive in test_online_data_update_sampling_weights_with_templated_eval_dataset: OnlineMixingDataset.log["count"] starts as an int literal in dataset.py before being overwritten with a list by update_sampling_weights(), so pylint's static inference can't see it's iterable by the time the test reads it. Wrapping in list() satisfies the checker without changing runtime behavior. Co-Authored-By: Claude Sonnet 5 Signed-off-by: Amal Joe --- .../src/fms_acceleration_odm/odm/dataset.py | 32 ++- .../src/fms_acceleration_odm/odm/reward.py | 194 +++++++++++------- .../tests/test_online_data.py | 3 +- 3 files changed, 140 insertions(+), 89 deletions(-) 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 66f2b1d4..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 @@ -513,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: + if state is None or self.reward_type.startswith(Reward.ENTROPY): return {} - if 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": [ @@ -528,21 +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] - } - if self.reward_type == Reward.LEARNABILITY: - return {} - if self.reward_type == Reward.VELOCITY: - return {} - if self.reward_type == Reward.COMBINED: - return { + }, + Reward.COMBINED: lambda: { "train_step": state.global_step, "total_steps": getattr(state, "max_steps", None), "beta": self.beta, - } - return {} + }, + } + 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 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 f533e50d..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 @@ -120,86 +120,138 @@ def compute_reward( 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] - log_probs = F.log_softmax(shift_logits, dim=-1) - probs = torch.exp(log_probs) + mask = batch["attention_mask"][:, 1:] - 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 * mask).sum(dim=-1) / mask.sum(dim=-1) + varentropy = (varentropy * mask).sum(dim=-1) / mask.sum(dim=-1) - entropy_last_token = entropy[:, -1] + max_entropy = torch.log( + torch.tensor(vocab_size, dtype=entropy.dtype, device=entropy.device) + ) - mask = batch["attention_mask"][:, 1:] + 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() + + +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." + ) + 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." ) - return GRADNORM_DATA["buffer"][current_category] - if reward_type == Reward.LEARNABILITY: - return _compute_learnability_reward(model, zero_shot_batch, few_shot_batch) - if reward_type == Reward.VELOCITY: - return _compute_velocity_reward(model, batch, current_category, total_categories) - if reward_type == Reward.COMBINED: - 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 - raise TypeError(f"Reward {reward_type} not supported") + 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: diff --git a/plugins/online-data-mixing/tests/test_online_data.py b/plugins/online-data-mixing/tests/test_online_data.py index e76d259b..daa7f0ee 100644 --- a/plugins/online-data-mixing/tests/test_online_data.py +++ b/plugins/online-data-mixing/tests/test_online_data.py @@ -170,6 +170,7 @@ class DummyState: dataset.update_sampling_weights(model, accelerator=None, 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 dataset.log["count"] + count > 0 for count in counts ), "expected every category to accumulate a nonzero count" From bc651dc34750e1156705d0efdfa86002f243c687 Mon Sep 17 00:00:00 2001 From: Amal Joe Date: Tue, 8 Sep 2026 22:08:03 +0530 Subject: [PATCH 3/3] fix: avoid hardcoded torch.device(0) in templated-eval-dataset test update_sampling_weights() falls back to torch.device(0) (cuda:0) when no accelerator is passed, so the test previously moved the model there too. That happens to resolve to mps:0 on Apple Silicon but crashes with "Found no NVIDIA driver" on the CPU-only CI runner, since torch.device(0) always means CUDA regardless of platform. Pass a minimal single-process accelerator stub (device=cpu, no-op prepare/reduce) instead, so the test exercises the same accelerator-driven code path update_sampling_weights() takes under real (single-process) Accelerate usage, without depending on any GPU/MPS backend being present. Co-Authored-By: Claude Sonnet 5 Signed-off-by: Amal Joe --- .../tests/test_online_data.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/plugins/online-data-mixing/tests/test_online_data.py b/plugins/online-data-mixing/tests/test_online_data.py index daa7f0ee..7a13235a 100644 --- a/plugins/online-data-mixing/tests/test_online_data.py +++ b/plugins/online-data-mixing/tests/test_online_data.py @@ -163,11 +163,20 @@ class DummyState: 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 torch.device(0) when no - # accelerator is given; match the model to whatever that resolves to here. - model = model.to(torch.device(0)) - dataset.update_sampling_weights(model, accelerator=None, state=DummyState()) + # 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"])