A context manager that enables stratified sampling, distributing the strata across all processes and gradient accumulation steps using settings from Hugging Face Accelerate.
(accelerator, disable=False)
| 296 | |
| 297 | @contextmanager |
| 298 | def enable_stratified_accelerate(accelerator, disable=False): |
| 299 | """A context manager that enables stratified sampling, distributing the strata across |
| 300 | all processes and gradient accumulation steps using settings from Hugging Face Accelerate.""" |
| 301 | try: |
| 302 | rank = accelerator.process_index |
| 303 | world_size = accelerator.num_processes |
| 304 | acc_steps = accelerator.gradient_state.num_steps |
| 305 | acc_step = accelerator.step % acc_steps |
| 306 | group = rank * acc_steps + acc_step |
| 307 | groups = world_size * acc_steps |
| 308 | with enable_stratified(group, groups, disable=disable): |
| 309 | yield |
| 310 | finally: |
| 311 | pass |
| 312 | |
| 313 | |
| 314 | def stratified_with_settings(shape, dtype=None, device=None): |
nothing calls this directly
no test coverage detected