From cab9b30788133f89b1d45e1d06421defbb50b2ff Mon Sep 17 00:00:00 2001 From: Adam Lee Date: Tue, 26 May 2026 23:45:12 -0700 Subject: [PATCH 1/2] [genrl]: keep repeated prompt samples on one rank Extracted from #1391. GenRL-Stack: 2/6 --- fastvideo/train/methods/rl/utils/data.py | 33 ++++++++++-------------- 1 file changed, 13 insertions(+), 20 deletions(-) diff --git a/fastvideo/train/methods/rl/utils/data.py b/fastvideo/train/methods/rl/utils/data.py index b0e0a787bd..db54011d40 100644 --- a/fastvideo/train/methods/rl/utils/data.py +++ b/fastvideo/train/methods/rl/utils/data.py @@ -154,12 +154,19 @@ def __init__( self.rank = rank self.seed = seed self.total_samples = num_replicas * batch_size + if self.batch_size % self.k != 0: + raise ValueError( + "batch_size must be divisible by k so each rank receives " + "whole prompt groups. Got " + f"batch_size={batch_size}, k={k}." + ) assert self.total_samples % self.k == 0, ( f"k cannot divide n*b: k={k} " f"num_replicas={num_replicas} " f"batch_size={batch_size}" ) self.m = self.total_samples // self.k + self.groups_per_rank = self.batch_size // self.k self.epoch = 0 def __iter__(self): @@ -169,28 +176,14 @@ def __iter__(self): indices = torch.randperm( len(self.dataset), generator=g )[: self.m].tolist() - repeated = [ - idx - for idx in indices + start = self.rank * self.groups_per_rank + end = start + self.groups_per_rank + rank_groups = indices[start:end] + yield [ + (self.epoch, idx) + for idx in rank_groups for _ in range(self.k) ] - shuffled_idx = torch.randperm( - len(repeated), generator=g - ).tolist() - shuffled = [ - repeated[i] for i in shuffled_idx - ] - per_card = [] - for i in range(self.num_replicas): - start = i * self.batch_size - end = start + self.batch_size - per_card.append( - [ - (self.epoch, idx) - for idx in shuffled[start:end] - ] - ) - yield per_card[self.rank] def set_epoch(self, epoch: int): self.epoch = epoch From 14ac93696b15c74ccd1c180fde3f78d96f39e036 Mon Sep 17 00:00:00 2001 From: Davids048 Date: Thu, 4 Jun 2026 19:07:11 +0000 Subject: [PATCH 2/2] [patch]: validate GenRL prompt sampler inputs --- fastvideo/train/methods/rl/utils/data.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/fastvideo/train/methods/rl/utils/data.py b/fastvideo/train/methods/rl/utils/data.py index db54011d40..00e2977171 100644 --- a/fastvideo/train/methods/rl/utils/data.py +++ b/fastvideo/train/methods/rl/utils/data.py @@ -149,11 +149,15 @@ def __init__( ): self.dataset = dataset self.batch_size = batch_size - self.k = k + self.k = k # Repeats/videos per prompt. self.num_replicas = num_replicas self.rank = rank self.seed = seed self.total_samples = num_replicas * batch_size + if self.k <= 0: + raise ValueError( + f"k must be a positive integer. Got k={k}." + ) if self.batch_size % self.k != 0: raise ValueError( "batch_size must be divisible by k so each rank receives " @@ -165,7 +169,13 @@ def __init__( f"num_replicas={num_replicas} " f"batch_size={batch_size}" ) - self.m = self.total_samples // self.k + self.m = self.total_samples // self.k # Unique prompts across ranks. + if len(self.dataset) < self.m: + raise ValueError( + "dataset must contain at least one prompt per global " + "prompt group. Got " + f"dataset_size={len(self.dataset)}, required={self.m}." + ) self.groups_per_rank = self.batch_size // self.k self.epoch = 0