Skip to content

Commit 9f9c8d5

Browse files
authored
[feat] GenRL: keep repeated prompt samples on one rank (#1401)
1 parent 28268ce commit 9f9c8d5

1 file changed

Lines changed: 25 additions & 22 deletions

File tree

  • fastvideo/train/methods/rl/utils

fastvideo/train/methods/rl/utils/data.py

Lines changed: 25 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -149,17 +149,34 @@ def __init__(
149149
):
150150
self.dataset = dataset
151151
self.batch_size = batch_size
152-
self.k = k
152+
self.k = k # Repeats/videos per prompt.
153153
self.num_replicas = num_replicas
154154
self.rank = rank
155155
self.seed = seed
156156
self.total_samples = num_replicas * batch_size
157+
if self.k <= 0:
158+
raise ValueError(
159+
f"k must be a positive integer. Got k={k}."
160+
)
161+
if self.batch_size % self.k != 0:
162+
raise ValueError(
163+
"batch_size must be divisible by k so each rank receives "
164+
"whole prompt groups. Got "
165+
f"batch_size={batch_size}, k={k}."
166+
)
157167
assert self.total_samples % self.k == 0, (
158168
f"k cannot divide n*b: k={k} "
159169
f"num_replicas={num_replicas} "
160170
f"batch_size={batch_size}"
161171
)
162-
self.m = self.total_samples // self.k
172+
self.m = self.total_samples // self.k # Unique prompts across ranks.
173+
if len(self.dataset) < self.m:
174+
raise ValueError(
175+
"dataset must contain at least one prompt per global "
176+
"prompt group. Got "
177+
f"dataset_size={len(self.dataset)}, required={self.m}."
178+
)
179+
self.groups_per_rank = self.batch_size // self.k
163180
self.epoch = 0
164181

165182
def __iter__(self):
@@ -169,28 +186,14 @@ def __iter__(self):
169186
indices = torch.randperm(
170187
len(self.dataset), generator=g
171188
)[: self.m].tolist()
172-
repeated = [
173-
idx
174-
for idx in indices
189+
start = self.rank * self.groups_per_rank
190+
end = start + self.groups_per_rank
191+
rank_groups = indices[start:end]
192+
yield [
193+
(self.epoch, idx)
194+
for idx in rank_groups
175195
for _ in range(self.k)
176196
]
177-
shuffled_idx = torch.randperm(
178-
len(repeated), generator=g
179-
).tolist()
180-
shuffled = [
181-
repeated[i] for i in shuffled_idx
182-
]
183-
per_card = []
184-
for i in range(self.num_replicas):
185-
start = i * self.batch_size
186-
end = start + self.batch_size
187-
per_card.append(
188-
[
189-
(self.epoch, idx)
190-
for idx in shuffled[start:end]
191-
]
192-
)
193-
yield per_card[self.rank]
194197

195198
def set_epoch(self, epoch: int):
196199
self.epoch = epoch

0 commit comments

Comments
 (0)