MCPcopy Create free account
hub / github.com/Meshcapade/difflocks / enable_stratified_accelerate

Function enable_stratified_accelerate

k_diffusion/utils.py:298–311  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

296
297@contextmanager
298def 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
314def stratified_with_settings(shape, dtype=None, device=None):

Callers

nothing calls this directly

Calls 1

enable_stratifiedFunction · 0.85

Tested by

no test coverage detected