Added initial work on two novel context consistency techs: NaiveReuse and ContextRef; with this commit, NaiveReuse is functional, but not mask_opt, ContextRef is totally unfunctional as will require appropriate rework of Advanced-ControlNet refcn code to be pushed out
This commit is contained in:
+11
-2
@@ -11,8 +11,10 @@ import comfy.samplers
|
||||
from comfy.model_base import BaseModel
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context_extras import ContextExtrasGroup
|
||||
from .utils_motion import get_sorted_list_via_attr
|
||||
|
||||
|
||||
class ContextFuseMethod:
|
||||
FLAT = "flat"
|
||||
PYRAMID = "pyramid"
|
||||
@@ -76,6 +78,7 @@ class ContextOptions:
|
||||
class ContextOptionsGroup:
|
||||
def __init__(self):
|
||||
self.contexts: list[ContextOptions] = []
|
||||
self.extras = ContextExtrasGroup()
|
||||
self._current_context: ContextOptions = None
|
||||
self._current_used_steps: int = 0
|
||||
self._current_index: int = 0
|
||||
@@ -121,9 +124,10 @@ class ContextOptionsGroup:
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.contexts) == 0
|
||||
|
||||
|
||||
def clone(self):
|
||||
cloned = ContextOptionsGroup()
|
||||
cloned.extras = self.extras.clone()
|
||||
for context in self.contexts:
|
||||
cloned.contexts.append(context)
|
||||
cloned._set_first_as_current()
|
||||
@@ -132,6 +136,11 @@ class ContextOptionsGroup:
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
for context in self.contexts:
|
||||
context.start_t = model.model_sampling.percent_to_sigma(context.start_percent)
|
||||
self.extras.initialize_timesteps(model)
|
||||
|
||||
def prepare_current(self, t: Tensor):
|
||||
self.prepare_current_context(t)
|
||||
self.extras.prepare_current(t)
|
||||
|
||||
def prepare_current_context(self, t: Tensor):
|
||||
curr_t: float = t[0]
|
||||
@@ -620,7 +629,7 @@ def generate_context_visualization(context_opts: ContextOptionsGroup, model: Mod
|
||||
|
||||
for i, t in enumerate(sigmas):
|
||||
# make context_opts reflect current step/sigma
|
||||
context_opts.prepare_current_context([t])
|
||||
context_opts.prepare_current([t])
|
||||
context_opts.step = start_step+i
|
||||
|
||||
# check if context should even be active in this case
|
||||
|
||||
@@ -0,0 +1,112 @@
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
|
||||
class ContextExtra:
|
||||
def __init__(self, start_percent: float, end_percent: float):
|
||||
# scheduling
|
||||
self.start_percent = float(start_percent)
|
||||
self.start_t = 999999999.9
|
||||
self.end_percent = float(end_percent)
|
||||
self.end_t = 0.0
|
||||
self.curr_t = 999999999.9
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
self.start_t = model.model_sampling.percent_to_sigma(self.start_percent)
|
||||
self.end_t = model.model_sampling.percent_to_sigma(self.end_percent)
|
||||
|
||||
def prepare_current(self, t: Tensor):
|
||||
self.curr_t = t[0]
|
||||
|
||||
def should_run(self):
|
||||
if self.curr_t > self.start_t or self.curr_t < self.end_t:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
################################
|
||||
# Context Ref
|
||||
class ContextRefParams:
|
||||
def __init__(self,
|
||||
attn_style_fidelity=0.0, attn_ref_weight=0.0, attn_atrength=0.0,
|
||||
adain_style_fidelity=0.0, adain_ref_weight=0.0, adain_strength=0.0):
|
||||
# attn1
|
||||
self.attn_style_fidelity = attn_style_fidelity
|
||||
self.attn_ref_weight = attn_ref_weight
|
||||
self.attn_strength = attn_atrength
|
||||
# adain
|
||||
self.adain_style_fidelity = adain_style_fidelity
|
||||
self.adain_ref_weight = adain_ref_weight
|
||||
self.adain_strength = adain_strength
|
||||
|
||||
|
||||
class ContextRef(ContextExtra):
|
||||
def __init__(self, start_percent: float, end_percent: float, params: ContextRefParams):
|
||||
super().__init__(start_percent=start_percent, end_percent=end_percent)
|
||||
self.params = params
|
||||
|
||||
def should_run(self):
|
||||
return super().should_run()
|
||||
#--------------------------------
|
||||
|
||||
|
||||
################################
|
||||
# NaiveReuse
|
||||
class NaiveReuse(ContextExtra):
|
||||
def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, mask_opt: Tensor=None):
|
||||
super().__init__(start_percent=start_percent, end_percent=end_percent)
|
||||
self.weighted_mean = weighted_mean
|
||||
self.mask_opt = mask_opt
|
||||
|
||||
def should_run(self):
|
||||
to_return = super().should_run()
|
||||
# if weighted_mean is 0.0, then reuse will take no effect anyway
|
||||
return to_return and self.weighted_mean > 0.0
|
||||
#--------------------------------
|
||||
|
||||
|
||||
class ContextExtrasGroup:
|
||||
def __init__(self):
|
||||
self.context_ref: ContextRef = None
|
||||
self.naive_reuse: NaiveReuse = None
|
||||
|
||||
def get_extras_list(self) -> list[ContextExtra]:
|
||||
extras_list = []
|
||||
if self.context_ref is not None:
|
||||
extras_list.append(self.context_ref)
|
||||
if self.naive_reuse is not None:
|
||||
extras_list.append(self.naive_reuse)
|
||||
return extras_list
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
for extra in self.get_extras_list():
|
||||
extra.initialize_timesteps(model)
|
||||
|
||||
def prepare_current(self, t: Tensor):
|
||||
for extra in self.get_extras_list():
|
||||
extra.prepare_current(t)
|
||||
|
||||
def should_run_context_ref(self):
|
||||
if not self.context_ref:
|
||||
return False
|
||||
return self.context_ref.should_run()
|
||||
|
||||
def should_run_naive_reuse(self):
|
||||
if not self.naive_reuse:
|
||||
return False
|
||||
return self.naive_reuse.should_run()
|
||||
|
||||
def add(self, extra: ContextExtra):
|
||||
if type(extra) == ContextRef:
|
||||
self.context_ref = extra
|
||||
elif type(extra) == NaiveReuse:
|
||||
self.naive_reuse = extra
|
||||
else:
|
||||
raise Exception(f"Unrecognized ContextExtras type: {type(extra)}")
|
||||
|
||||
def clone(self):
|
||||
cloned = ContextExtrasGroup()
|
||||
cloned.context_ref = self.context_ref
|
||||
cloned.naive_reuse = self.naive_reuse
|
||||
return cloned
|
||||
+20
-3
@@ -28,7 +28,8 @@ from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, Sampl
|
||||
from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode, SigmaScheduleToSigmasNode)
|
||||
from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode,
|
||||
StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode,
|
||||
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom)
|
||||
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom,
|
||||
SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef)
|
||||
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
|
||||
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
|
||||
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
|
||||
@@ -54,13 +55,15 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ADE_MultivalDynamicFloatInput": MultivalDynamicFloatInputNode,
|
||||
"ADE_MultivalScaledMask": MultivalScaledMaskNode,
|
||||
"ADE_MultivalConvertToMask": MultivalConvertToMaskNode,
|
||||
###############################################################################
|
||||
#------------------------------------------------------------------------------
|
||||
# Context Opts
|
||||
"ADE_StandardStaticContextOptions": StandardStaticContextOptionsNode,
|
||||
"ADE_StandardUniformContextOptions": StandardUniformContextOptionsNode,
|
||||
"ADE_LoopedUniformContextOptions": LoopedUniformContextOptionsNode,
|
||||
"ADE_ViewsOnlyContextOptions": ViewAsContextOptionsNode,
|
||||
"ADE_BatchedContextOptions": BatchedContextOptionsNode,
|
||||
"ADE_AnimateDiffUniformContextOptions": LegacyLoopedUniformContextOptionsNode, # Legacy
|
||||
"ADE_AnimateDiffUniformContextOptions": LegacyLoopedUniformContextOptionsNode, # Legacy/Deprecated
|
||||
"ADE_VisualizeContextOptionsK": VisualizeContextOptionsK,
|
||||
"ADE_VisualizeContextOptionsKAdv": VisualizeContextOptionsKAdv,
|
||||
"ADE_VisualizeContextOptionsSCustom": VisualizeContextOptionsSCustom,
|
||||
@@ -68,6 +71,12 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ADE_StandardStaticViewOptions": StandardStaticViewOptionsNode,
|
||||
"ADE_StandardUniformViewOptions": StandardUniformViewOptionsNode,
|
||||
"ADE_LoopedUniformViewOptions": LoopedUniformViewOptionsNode,
|
||||
# Context Extras
|
||||
"ADE_ContextExtras_Set": SetContextExtrasOnContextOptions,
|
||||
"ADE_ContextExtras_ContextRef": ContextExtras_ContextRef,
|
||||
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
|
||||
#------------------------------------------------------------------------------
|
||||
###############################################################################
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": IterationOptionsNode,
|
||||
"ADE_IterationOptsFreeInit": FreeInitOptionsNode,
|
||||
@@ -184,13 +193,15 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ADE_MultivalDynamicFloatInput": "Multival [Float List] 🎭🅐🅓",
|
||||
"ADE_MultivalScaledMask": "Multival Scaled Mask 🎭🅐🅓",
|
||||
"ADE_MultivalConvertToMask": "Multival to Mask 🎭🅐🅓",
|
||||
###############################################################################
|
||||
#------------------------------------------------------------------------------
|
||||
# Context Opts
|
||||
"ADE_StandardStaticContextOptions": "Context Options◆Standard Static 🎭🅐🅓",
|
||||
"ADE_StandardUniformContextOptions": "Context Options◆Standard Uniform 🎭🅐🅓",
|
||||
"ADE_LoopedUniformContextOptions": "Context Options◆Looped Uniform 🎭🅐🅓",
|
||||
"ADE_ViewsOnlyContextOptions": "Context Options◆Views Only [VRAM⇈] 🎭🅐🅓",
|
||||
"ADE_BatchedContextOptions": "Context Options◆Batched [Non-AD] 🎭🅐🅓",
|
||||
"ADE_AnimateDiffUniformContextOptions": "Context Options◆Looped Uniform 🎭🅐🅓", # Legacy
|
||||
"ADE_AnimateDiffUniformContextOptions": "Context Options◆Looped Uniform 🎭🅐🅓", # Legacy/Deprecated
|
||||
"ADE_VisualizeContextOptionsK": "Visualize Context Options (K.) 🎭🅐🅓",
|
||||
"ADE_VisualizeContextOptionsKAdv": "Visualize Context Options (K.Adv.) 🎭🅐🅓",
|
||||
"ADE_VisualizeContextOptionsSCustom": "Visualize Context Options (S.Cus.) 🎭🅐🅓",
|
||||
@@ -198,6 +209,12 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ADE_StandardStaticViewOptions": "View Options◆Standard Static 🎭🅐🅓",
|
||||
"ADE_StandardUniformViewOptions": "View Options◆Standard Uniform 🎭🅐🅓",
|
||||
"ADE_LoopedUniformViewOptions": "View Options◆Looped Uniform 🎭🅐🅓",
|
||||
# Context Extras
|
||||
"ADE_ContextExtras_Set": "Set Context Extras 🎭🅐🅓",
|
||||
"ADE_ContextExtras_ContextRef": "Context Extras◆ContextRef 🎭🅐🅓",
|
||||
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
|
||||
#------------------------------------------------------------------------------
|
||||
###############################################################################
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": "Default Iteration Options 🎭🅐🅓",
|
||||
"ADE_IterationOptsFreeInit": "FreeInit Iteration Options 🎭🅐🅓",
|
||||
|
||||
@@ -6,6 +6,7 @@ from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import (ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules,
|
||||
generate_context_visualization)
|
||||
from .context_extras import ContextExtrasGroup, ContextRef, ContextRefParams, NaiveReuse
|
||||
from .utils_model import BIGMAX, MAX_RESOLUTION
|
||||
|
||||
|
||||
@@ -440,3 +441,90 @@ class VisualizeContextOptionsSCustom:
|
||||
images = generate_context_visualization(context_opts=context_opts, model=model, width=visual_width, video_length=latents_length,
|
||||
sigmas=sigmas)
|
||||
return (images,)
|
||||
|
||||
|
||||
#########################
|
||||
# Context Extras
|
||||
class SetContextExtrasOnContextOptions:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"context_opts": ("CONTEXT_OPTIONS",),
|
||||
"context_extras": ("CONTEXT_EXTRAS",),
|
||||
},
|
||||
"optional": {
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras"
|
||||
FUNCTION = "set_context_extras"
|
||||
|
||||
def set_context_extras(self, context_opts: ContextOptionsGroup, context_extras: ContextExtrasGroup):
|
||||
context_opts = context_opts.clone()
|
||||
context_opts.extras = context_extras.clone()
|
||||
return (context_opts,)
|
||||
|
||||
|
||||
class ContextExtras_NaiveReuse:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
},
|
||||
"optional": {
|
||||
"prev_extras": ("CONTEXT_EXTRAS",),
|
||||
"mask_opt": ("MASK",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"weighted_mean": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 55}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_EXTRAS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras"
|
||||
FUNCTION = "create_context_extra"
|
||||
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None):
|
||||
if prev_extras is None:
|
||||
prev_extras = prev_extras = ContextExtrasGroup()
|
||||
prev_extras = prev_extras.clone()
|
||||
# create extra
|
||||
naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, mask_opt=mask_opt)
|
||||
prev_extras.add(naive_reuse)
|
||||
return (prev_extras,)
|
||||
|
||||
|
||||
class ContextExtras_ContextRef:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
},
|
||||
"optional": {
|
||||
"prev_extras": ("CONTEXT_EXTRAS",),
|
||||
"mask_opt": ("MASK",),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"autosize": ("ADEAUTOSIZE", {"padding": 55}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_EXTRAS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras"
|
||||
FUNCTION = "create_context_extra"
|
||||
|
||||
def create_context_extra(self, start_percent=0.0, end_percent=0.1, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None):
|
||||
if prev_extras is None:
|
||||
prev_extras = prev_extras = ContextExtrasGroup()
|
||||
prev_extras = prev_extras.clone()
|
||||
# create extra
|
||||
# TODO: make customizable, and allow mask input
|
||||
params = ContextRefParams(attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_atrength=1.0)
|
||||
context_ref = ContextRef(start_percent=start_percent, end_percent=end_percent, params=params)
|
||||
prev_extras.add(context_ref)
|
||||
return (prev_extras,)
|
||||
|
||||
+58
-4
@@ -71,7 +71,7 @@ class AnimateDiffHelper_GlobalState:
|
||||
if self.motion_models is not None:
|
||||
self.motion_models.prepare_current_keyframe(x=x, t=timestep)
|
||||
if self.params.context_options is not None:
|
||||
self.params.context_options.prepare_current_context(t=timestep)
|
||||
self.params.context_options.prepare_current(t=timestep)
|
||||
if self.sample_settings.custom_cfg is not None:
|
||||
self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep)
|
||||
|
||||
@@ -617,8 +617,10 @@ def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond,
|
||||
|
||||
# add AD/evolved-sampling params to model_options (transformer_options)
|
||||
model_options = model_options.copy()
|
||||
if "tranformer_options" not in model_options:
|
||||
model_options["tranformer_options"] = {}
|
||||
if "transformer_options" not in model_options:
|
||||
model_options["transformer_options"] = {}
|
||||
else:
|
||||
model_options["transformer_options"] = model_options["transformer_options"].copy()
|
||||
model_options["transformer_options"]["ad_params"] = ADGS.create_exposed_params()
|
||||
|
||||
if not ADGS.is_using_sliding_context():
|
||||
@@ -798,7 +800,25 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
|
||||
counts_final = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds]
|
||||
biases_final = [([0.0] * x_in.shape[0]) for _ in conds]
|
||||
|
||||
# perform calc_conds_batch per context window
|
||||
CONTEXTREF_ATTN_MACHINE_STATE = "contextref_attn_machine_state"
|
||||
CONTEXTREF_ADAIN_MACHINE_STATE = "contextref_adain_machine_state"
|
||||
#context_ref_steps = [0,1,2,3,4,5]#[0]
|
||||
context_ref = False
|
||||
first_context = False
|
||||
if ADGS.params.context_options.extras.should_run_context_ref():
|
||||
context_ref = True
|
||||
first_context = True
|
||||
|
||||
#naive_steps = [0,1]#[0,1,2]#[0,1,2,3]
|
||||
naive_counts_mult = 50
|
||||
naive_init = False
|
||||
cached_naive_conds = None
|
||||
cached_naive_ctx_idxs = None
|
||||
if ADGS.params.context_options.extras.should_run_naive_reuse():
|
||||
cached_naive_conds = [torch.zeros_like(x_in) for _ in conds]
|
||||
#cached_naive_counts = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds]
|
||||
naive_init = True
|
||||
# perform calc_conds_batch per context window
|
||||
for ctx_idxs in context_windows:
|
||||
ADGS.params.sub_idxs = ctx_idxs
|
||||
if ADGS.motion_models is not None:
|
||||
@@ -817,6 +837,18 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
|
||||
sub_timestep = timestep[full_idxs]
|
||||
sub_conds = [get_resized_cond(cond, full_idxs, len(ctx_idxs)) for cond in conds]
|
||||
|
||||
if context_ref:
|
||||
if first_context:
|
||||
first_context = False
|
||||
model_options["transformer_options"][CONTEXTREF_ATTN_MACHINE_STATE] = "write"
|
||||
model_options["transformer_options"][CONTEXTREF_ADAIN_MACHINE_STATE] = "write"
|
||||
else:
|
||||
model_options["transformer_options"][CONTEXTREF_ATTN_MACHINE_STATE] = "read"
|
||||
model_options["transformer_options"][CONTEXTREF_ADAIN_MACHINE_STATE] = "read"
|
||||
else:
|
||||
model_options["transformer_options"][CONTEXTREF_ATTN_MACHINE_STATE] = "off"
|
||||
model_options["transformer_options"][CONTEXTREF_ADAIN_MACHINE_STATE] = "off"
|
||||
|
||||
sub_conds_out = calc_cond_uncond_batch_wrapper(model, sub_conds, sub_x, sub_timestep, model_options)
|
||||
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
@@ -841,12 +873,34 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
|
||||
for i in range(len(sub_conds_out)):
|
||||
conds_final[i][full_idxs] += sub_conds_out[i] * weights_tensor
|
||||
counts_final[i][full_idxs] += weights_tensor
|
||||
if naive_init:
|
||||
cached_naive_ctx_idxs = ctx_idxs
|
||||
for i in range(len(sub_conds)):
|
||||
cached_naive_conds[i][full_idxs] = conds_final[i][full_idxs] / counts_final[i][full_idxs]
|
||||
#cached_naive_conds[i][full_idxs] = conds_final[i][full_idxs]
|
||||
#cached_naive_counts[i][full_idxs] = counts_final[i][full_idxs]
|
||||
naive_init = False
|
||||
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
# already normalized, so return as is
|
||||
del counts_final
|
||||
return conds_final
|
||||
else:
|
||||
if cached_naive_conds is not None:
|
||||
#start_idx = cached_naive_ctx_idxs[-1] + 1
|
||||
start_idx = cached_naive_ctx_idxs[0]
|
||||
for z in range(start_idx, ADGS.params.full_length, len(cached_naive_ctx_idxs)):
|
||||
for i in range(len(cached_naive_conds)):
|
||||
new_ctx_idxs = [zz for zz in list(range(z, z+len(cached_naive_ctx_idxs))) if zz < ADGS.params.full_length]
|
||||
# make sure when getting cached_naive idxs, they are adjusted for actual length leftover length
|
||||
adjusted_cnaive_ctx_idxs = cached_naive_ctx_idxs[:len(new_ctx_idxs)]
|
||||
weighted_mean = ADGS.params.context_options.extras.naive_reuse.weighted_mean
|
||||
conds_final[i][new_ctx_idxs] = (weighted_mean * (cached_naive_conds[i][adjusted_cnaive_ctx_idxs]*counts_final[i][new_ctx_idxs])) + ((1.-weighted_mean) * conds_final[i][new_ctx_idxs])
|
||||
#conds_final[i][new_idxs] += (cached_naive_conds[i][cached_naive_full_idxs] / cached_naive_counts[i][cached_naive_full_idxs]) * counts
|
||||
#counts = counts_final[i][new_idxs] * naive_counts_mult# / 2
|
||||
#conds_final[i][new_idxs] += (cached_naive_conds[i][cached_naive_full_idxs] / cached_naive_counts[i][cached_naive_full_idxs]) * counts
|
||||
#counts_final[i][new_idxs] += counts# * 10#counts_final[i][full_idxs]
|
||||
del cached_naive_conds
|
||||
# normalize conds via division by context usage counts
|
||||
for i in range(len(conds_final)):
|
||||
conds_final[i] /= counts_final[i]
|
||||
|
||||
Reference in New Issue
Block a user