Merge pull request #28 from AustinMroz/main

Possible fixes to uniform context window generator - available as "uniform v2" context scheduler
This commit is contained in:
Jedrzej Kosinski
2023-09-19 16:40:06 -05:00
committed by GitHub
+49 -1
View File
@@ -7,8 +7,9 @@ import numpy as np
class ContextSchedules:
UNIFORM = "uniform"
UNIFORM_CONSTANT = "uniform_constant"
UNIFORM_V2 = "uniform v2"
CONTEXT_SCHEDULE_LIST = [UNIFORM]
CONTEXT_SCHEDULE_LIST = [UNIFORM, UNIFORM_V2]
# Returns fraction that has denominator that is a power of 2
@@ -53,6 +54,51 @@ def uniform(
):
yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)]
def uniform_v2(
step: int = ...,
num_steps: Optional[int] = None,
num_frames: int = ...,
context_size: Optional[int] = None,
context_stride: int = 3,
context_overlap: int = 4,
closed_loop: bool = True,
print_final: bool = False,
):
if num_frames <= context_size:
yield list(range(num_frames))
return
context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1)
pad = int(round(num_frames * ordered_halving(step, print_final)))
for context_step in 1 << np.arange(context_stride):
j_initial = int(ordered_halving(step) * context_step) + pad
for j in range(
j_initial,
num_frames + pad - context_overlap,
(context_size * context_step - context_overlap),
):
if context_size * context_step > num_frames:
# On the final context_step,
# ensure no frame appears in the window twice
yield [e % num_frames for e in range(j, j + num_frames, context_step)]
continue
j = j % num_frames
if j > (j + context_size * context_step) % num_frames and not closed_loop:
yield [e for e in range(j, num_frames, context_step)]
j_stop = (j + context_size * context_step) % num_frames
# When ((num_frames % (context_size - context_overlap)+context_overlap) % context_size != 0,
# This can cause 'superflous' runs where all frames in
# a context window have already been processed during
# the first context window of this stride and step.
# While the following commented if should prevent this,
# I believe leaving it in is more correct as it maintains
# the total conditional passes per frame over a large total steps
# if j_stop > context_overlap:
yield [e for e in range(0, j_stop, context_step)]
continue
yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)]
def uniform_constant(
step: int = ...,
@@ -103,6 +149,8 @@ def get_context_scheduler(name: str) -> Callable:
return uniform
case ContextSchedules.UNIFORM_CONSTANT:
return uniform_constant
case ContextSchedules.UNIFORM_V2:
return uniform_v2
case _:
raise ValueError(f"Unknown context_overlap policy {name}")