adaptive decay as an overridable param (+1 squashed commits)

Squashed commits:

[d94df7843] adaptive decay as an overridable param
This commit is contained in:
Concedo
2025-12-28 11:13:18 +08:00
parent 1051313cb2
commit 27261bfc26
4 changed files with 12 additions and 8 deletions
+4 -1
View File
@@ -266,6 +266,7 @@ class generation_inputs(ctypes.Structure):
("smoothing_factor", ctypes.c_float),
("smoothing_curve", ctypes.c_float),
("adaptive_target", ctypes.c_float),
("adaptive_decay", ctypes.c_float),
("dry_multiplier", ctypes.c_float),
("dry_base", ctypes.c_float),
("dry_allowed_length", ctypes.c_int),
@@ -1605,8 +1606,9 @@ def generate(genparams, stream_flag=False):
smoothing_factor = tryparsefloat(genparams.get('smoothing_factor', 0.0),0.0)
smoothing_curve = tryparsefloat(genparams.get('smoothing_curve', 1.0),1.0)
adaptive_target = tryparsefloat(genparams.get('adaptive_target', -1.0),-1.0)
adaptive_decay = tryparsefloat(genparams.get('adaptive_decay', 0.9),0.9)
if adaptive_target>0 and min_p<=0 and top_p>=1.0: #adaptive p sampler requires a truncation sampler first, force a tiny min-p
min_p = 0.01
min_p = 0.002
logit_biases = genparams.get('logit_bias', {})
render_special = genparams.get('render_special', False)
banned_strings = genparams.get('banned_strings', []) # SillyTavern uses that name
@@ -1671,6 +1673,7 @@ def generate(genparams, stream_flag=False):
inputs.smoothing_factor = smoothing_factor
inputs.smoothing_curve = smoothing_curve
inputs.adaptive_target = adaptive_target
inputs.adaptive_decay = adaptive_decay
inputs.grammar = grammar.encode("UTF-8")
inputs.grammar_retain_state = grammar_retain_state
inputs.allow_eos_token = not ban_eos_token