power law sampler added

This commit is contained in:
Concedo
2025-12-27 09:46:06 +08:00
parent dfa1b72d2f
commit 91d8863f18
5 changed files with 987 additions and 110 deletions
+7 -1
View File
@@ -265,6 +265,7 @@ class generation_inputs(ctypes.Structure):
("dynatemp_exponent", ctypes.c_float),
("smoothing_factor", ctypes.c_float),
("smoothing_curve", ctypes.c_float),
("power_law_target", ctypes.c_float),
("dry_multiplier", ctypes.c_float),
("dry_base", ctypes.c_float),
("dry_allowed_length", ctypes.c_int),
@@ -1603,6 +1604,9 @@ def generate(genparams, stream_flag=False):
dynatemp_exponent = tryparsefloat(genparams.get('dynatemp_exponent', 1.0),1.0)
smoothing_factor = tryparsefloat(genparams.get('smoothing_factor', 0.0),0.0)
smoothing_curve = tryparsefloat(genparams.get('smoothing_curve', 1.0),1.0)
power_law_target = tryparsefloat(genparams.get('power_law_target', -1.0),-1.0)
if power_law_target>0 and min_p<=0 and top_p>=1.0: #power law sampler requires a truncation sampler first, force a tiny min-p
min_p = 0.01
logit_biases = genparams.get('logit_bias', {})
render_special = genparams.get('render_special', False)
banned_strings = genparams.get('banned_strings', []) # SillyTavern uses that name
@@ -1666,6 +1670,7 @@ def generate(genparams, stream_flag=False):
inputs.dynatemp_exponent = dynatemp_exponent
inputs.smoothing_factor = smoothing_factor
inputs.smoothing_curve = smoothing_curve
inputs.power_law_target = power_law_target
inputs.grammar = grammar.encode("UTF-8")
inputs.grammar_retain_state = grammar_retain_state
inputs.allow_eos_token = not ban_eos_token
@@ -2814,7 +2819,8 @@ ws ::= | " " | "\n" [ \t]{0,20}
default_adapter = {} if chatcompl_adapter is None else chatcompl_adapter
adapter_obj = genparams.get('adapter', default_adapter)
default_max_tok = (adapter_obj.get("max_length", args.defaultgenamt) if (api_format==4 or api_format==7) else args.defaultgenamt)
genparams["max_length"] = tryparseint(genparams.get('max_tokens', genparams.get('max_completion_tokens', default_max_tok)),default_max_tok)
oaiml = tryparseint(genparams.get('max_tokens', genparams.get('max_completion_tokens', default_max_tok)),default_max_tok)
genparams["max_length"] = genparams.get('max_length', oaiml)
if genparams["max_length"] <= 0:
genparams["max_length"] = default_max_tok
presence_penalty = genparams.get('presence_penalty', genparams.get('frequency_penalty', 0.0))