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:
+49
-1
@@ -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}")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user