44
55from __future__ import annotations
66
7+ import hashlib
78import time
89from collections .abc import Callable
910from typing import Any
@@ -31,8 +32,15 @@ def create_generator(
3132 """Create deterministic generators seeded by prompt."""
3233 generators = []
3334 for prompt in prompts :
35+ prompt_seed = int .from_bytes (
36+ hashlib .blake2b (
37+ prompt .encode ("utf-8" ),
38+ digest_size = 8 ,
39+ ).digest (),
40+ "big" ,
41+ )
3442 g = torch .Generator (device = device )
35- g .manual_seed (base_seed + hash ( prompt ) % (2 ** 31 ))
43+ g .manual_seed (base_seed + prompt_seed % (2 ** 31 ))
3644 generators .append (g )
3745 return generators
3846
@@ -72,13 +80,13 @@ def sample_epoch(
7280 ref_transformer : torch .nn .Module | None = None ,
7381 lora_model : Any | None = None ,
7482 tracker : Any | None = None ,
83+ async_reward_scoring : bool = True ,
7584) -> tuple [
7685 list [dict [str , Any ]],
7786 list [torch .Tensor ],
7887 list [list [str ]],
7988]:
80- """Run one sampling epoch: generate videos, compute
81- rewards asynchronously.
89+ """Run one sampling epoch: generate videos and compute rewards.
8290
8391 Returns:
8492 Tuple of (samples, all_videos, all_prompts):
@@ -131,7 +139,7 @@ def sample_epoch(
131139 if same_latent :
132140 gen = create_generator (
133141 prompts ,
134- base_seed = epoch * SEED_EPOCH_STRIDE + i ,
142+ base_seed = seed + epoch * SEED_EPOCH_STRIDE + i ,
135143 device = device ,
136144 )
137145 else :
@@ -188,19 +196,25 @@ def sample_epoch(
188196 .repeat (sample_batch_size , 1 )
189197 )
190198
199+ videos_cpu = videos .detach ().cpu ()
200+
191201 # Collect decoded videos and prompts for logging.
192- all_videos .append (videos )
202+ all_videos .append (videos_cpu )
193203 all_prompts .append (list (prompts ))
194204
195- # Async reward computation.
196- rewards_future = executor .submit (
197- reward_fn ,
198- videos ,
199- prompts ,
200- prompt_metadata ,
201- True ,
202- )
203- time .sleep (0 )
205+ if async_reward_scoring :
206+ rewards = executor .submit (
207+ reward_fn ,
208+ videos_cpu ,
209+ prompts ,
210+ prompt_metadata ,
211+ True ,
212+ )
213+ time .sleep (0 )
214+ else :
215+ rewards = (videos_cpu , list (prompts ), prompt_metadata )
216+
217+ del videos
204218
205219 logger .info (
206220 "[sample_epoch] batch %d/%d: "
@@ -225,15 +239,22 @@ def sample_epoch(
225239 "next_latents" : latents [:, 1 :],
226240 "log_probs" : log_probs ,
227241 "kl" : kl ,
228- "rewards" : rewards_future ,
242+ "rewards" : rewards ,
229243 }
230244 )
231245
232246 # Wait for all rewards.
233247 torch .cuda .synchronize ()
234248 _t_reward_wait = time .perf_counter ()
235249 for sample in samples :
236- rewards , _ = sample ["rewards" ].result ()
250+ if async_reward_scoring :
251+ rewards , _ = sample ["rewards" ].result ()
252+ else :
253+ videos_cpu , prompts , prompt_metadata = sample ["rewards" ]
254+ torch .cuda .empty_cache ()
255+ rewards , _ = reward_fn (
256+ videos_cpu , prompts , prompt_metadata , True
257+ )
237258 sample ["rewards" ] = {
238259 key : torch .as_tensor (value , device = device ).float ()
239260 for key , value in rewards .items ()
0 commit comments