| 76 | |
| 77 | |
| 78 | def flow_grpo_step( |
| 79 | model_output: torch.Tensor, |
| 80 | latents: torch.Tensor, |
| 81 | eta: float, |
| 82 | sigmas: torch.Tensor, |
| 83 | index: int, |
| 84 | prev_sample: torch.Tensor, |
| 85 | generator: Optional[torch.Generator] = None, |
| 86 | ): |
| 87 | device = model_output.device |
| 88 | sigma = sigmas[index].to(device) |
| 89 | sigma_prev = sigmas[index + 1].to(device) |
| 90 | sigma_max = sigmas[1].item() |
| 91 | dt = sigma_prev - sigma # neg dt |
| 92 | |
| 93 | pred_original_sample = latents - sigma * model_output |
| 94 | |
| 95 | std_dev_t = torch.sqrt(sigma / (1 - torch.where(sigma == 1, sigma_max, sigma))) * eta |
| 96 | |
| 97 | if prev_sample is not None and generator is not None: |
| 98 | raise ValueError( |
| 99 | "Cannot pass both generator and prev_sample. Please make sure that either `generator` or" |
| 100 | " `prev_sample` stays `None`." |
| 101 | ) |
| 102 | |
| 103 | prev_sample_mean = ( |
| 104 | latents * (1 + std_dev_t**2 / (2 * sigma) * dt) |
| 105 | + model_output * (1 + std_dev_t**2 * (1 - sigma) / (2 * sigma)) * dt |
| 106 | ) |
| 107 | |
| 108 | if prev_sample is None: |
| 109 | variance_noise = randn_tensor(model_output.shape, generator=generator, device=device, dtype=model_output.dtype) |
| 110 | prev_sample = prev_sample_mean + std_dev_t * torch.sqrt(-1 * dt) * variance_noise |
| 111 | |
| 112 | log_prob = ( |
| 113 | -((prev_sample.detach() - prev_sample_mean) ** 2) / (2 * ((std_dev_t * torch.sqrt(-1 * dt)) ** 2)) |
| 114 | - torch.log(std_dev_t * torch.sqrt(-1 * dt)) |
| 115 | - torch.log(torch.sqrt(2 * torch.as_tensor(math.pi))) |
| 116 | ) |
| 117 | |
| 118 | # mean along all but batch dimension |
| 119 | log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim))) |
| 120 | |
| 121 | return prev_sample, pred_original_sample, log_prob |
| 122 | |
| 123 | |
| 124 | def dance_grpo_step( |