decider LoRA fine-tune
Browse files- .gitattributes +1 -0
- chat_template.jinja +154 -0
- config.json +83 -0
- decider/__init__.py +0 -0
- decider/__pycache__/__init__.cpython-311.pyc +0 -0
- decider/__pycache__/infer.cpython-311.pyc +0 -0
- decider/__pycache__/model.cpython-311.pyc +0 -0
- decider/__pycache__/prompt.cpython-311.pyc +0 -0
- decider/__pycache__/systemone.cpython-311.pyc +0 -0
- decider/__pycache__/temperature.cpython-311.pyc +0 -0
- decider/batching.py +95 -0
- decider/calibrate.py +222 -0
- decider/engine.py +216 -0
- decider/engine_v2.py +213 -0
- decider/fp8.py +54 -0
- decider/infer.py +357 -0
- decider/metrics.py +37 -0
- decider/model.py +49 -0
- decider/mps_moe.py +102 -0
- decider/mps_ops.py +268 -0
- decider/prompt.py +287 -0
- decider/prompt_fast.py +90 -0
- decider/schema_engine.py +150 -0
- decider/serve.py +562 -0
- decider/shared_prefix.py +200 -0
- decider/systemone.py +199 -0
- decider/temperature.py +144 -0
- decider_config.json +18 -0
- eval_report.json +331 -0
- export.json +20 -0
- generation_config.json +6 -0
- model.safetensors +3 -0
- tokenizer.json +3 -0
- tokenizer_config.json +32 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{%- set image_count = namespace(value=0) %}
|
| 2 |
+
{%- set video_count = namespace(value=0) %}
|
| 3 |
+
{%- macro render_content(content, do_vision_count, is_system_content=false) %}
|
| 4 |
+
{%- if content is string %}
|
| 5 |
+
{{- content }}
|
| 6 |
+
{%- elif content is iterable and content is not mapping %}
|
| 7 |
+
{%- for item in content %}
|
| 8 |
+
{%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
|
| 9 |
+
{%- if is_system_content %}
|
| 10 |
+
{{- raise_exception('System message cannot contain images.') }}
|
| 11 |
+
{%- endif %}
|
| 12 |
+
{%- if do_vision_count %}
|
| 13 |
+
{%- set image_count.value = image_count.value + 1 %}
|
| 14 |
+
{%- endif %}
|
| 15 |
+
{%- if add_vision_id %}
|
| 16 |
+
{{- 'Picture ' ~ image_count.value ~ ': ' }}
|
| 17 |
+
{%- endif %}
|
| 18 |
+
{{- '<|vision_start|><|image_pad|><|vision_end|>' }}
|
| 19 |
+
{%- elif 'video' in item or item.type == 'video' %}
|
| 20 |
+
{%- if is_system_content %}
|
| 21 |
+
{{- raise_exception('System message cannot contain videos.') }}
|
| 22 |
+
{%- endif %}
|
| 23 |
+
{%- if do_vision_count %}
|
| 24 |
+
{%- set video_count.value = video_count.value + 1 %}
|
| 25 |
+
{%- endif %}
|
| 26 |
+
{%- if add_vision_id %}
|
| 27 |
+
{{- 'Video ' ~ video_count.value ~ ': ' }}
|
| 28 |
+
{%- endif %}
|
| 29 |
+
{{- '<|vision_start|><|video_pad|><|vision_end|>' }}
|
| 30 |
+
{%- elif 'text' in item %}
|
| 31 |
+
{{- item.text }}
|
| 32 |
+
{%- else %}
|
| 33 |
+
{{- raise_exception('Unexpected item type in content.') }}
|
| 34 |
+
{%- endif %}
|
| 35 |
+
{%- endfor %}
|
| 36 |
+
{%- elif content is none or content is undefined %}
|
| 37 |
+
{{- '' }}
|
| 38 |
+
{%- else %}
|
| 39 |
+
{{- raise_exception('Unexpected content type.') }}
|
| 40 |
+
{%- endif %}
|
| 41 |
+
{%- endmacro %}
|
| 42 |
+
{%- if not messages %}
|
| 43 |
+
{{- raise_exception('No messages provided.') }}
|
| 44 |
+
{%- endif %}
|
| 45 |
+
{%- if tools and tools is iterable and tools is not mapping %}
|
| 46 |
+
{{- '<|im_start|>system\n' }}
|
| 47 |
+
{{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
|
| 48 |
+
{%- for tool in tools %}
|
| 49 |
+
{{- "\n" }}
|
| 50 |
+
{{- tool | tojson }}
|
| 51 |
+
{%- endfor %}
|
| 52 |
+
{{- "\n</tools>" }}
|
| 53 |
+
{{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
|
| 54 |
+
{%- if messages[0].role == 'system' %}
|
| 55 |
+
{%- set content = render_content(messages[0].content, false, true)|trim %}
|
| 56 |
+
{%- if content %}
|
| 57 |
+
{{- '\n\n' + content }}
|
| 58 |
+
{%- endif %}
|
| 59 |
+
{%- endif %}
|
| 60 |
+
{{- '<|im_end|>\n' }}
|
| 61 |
+
{%- else %}
|
| 62 |
+
{%- if messages[0].role == 'system' %}
|
| 63 |
+
{%- set content = render_content(messages[0].content, false, true)|trim %}
|
| 64 |
+
{{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
|
| 65 |
+
{%- endif %}
|
| 66 |
+
{%- endif %}
|
| 67 |
+
{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
|
| 68 |
+
{%- for message in messages[::-1] %}
|
| 69 |
+
{%- set index = (messages|length - 1) - loop.index0 %}
|
| 70 |
+
{%- if ns.multi_step_tool and message.role == "user" %}
|
| 71 |
+
{%- set content = render_content(message.content, false)|trim %}
|
| 72 |
+
{%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
|
| 73 |
+
{%- set ns.multi_step_tool = false %}
|
| 74 |
+
{%- set ns.last_query_index = index %}
|
| 75 |
+
{%- endif %}
|
| 76 |
+
{%- endif %}
|
| 77 |
+
{%- endfor %}
|
| 78 |
+
{%- if ns.multi_step_tool %}
|
| 79 |
+
{{- raise_exception('No user query found in messages.') }}
|
| 80 |
+
{%- endif %}
|
| 81 |
+
{%- for message in messages %}
|
| 82 |
+
{%- set content = render_content(message.content, true)|trim %}
|
| 83 |
+
{%- if message.role == "system" %}
|
| 84 |
+
{%- if not loop.first %}
|
| 85 |
+
{{- raise_exception('System message must be at the beginning.') }}
|
| 86 |
+
{%- endif %}
|
| 87 |
+
{%- elif message.role == "user" %}
|
| 88 |
+
{{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
|
| 89 |
+
{%- elif message.role == "assistant" %}
|
| 90 |
+
{%- set reasoning_content = '' %}
|
| 91 |
+
{%- if message.reasoning_content is string %}
|
| 92 |
+
{%- set reasoning_content = message.reasoning_content %}
|
| 93 |
+
{%- else %}
|
| 94 |
+
{%- if '</think>' in content %}
|
| 95 |
+
{%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
|
| 96 |
+
{%- set content = content.split('</think>')[-1].lstrip('\n') %}
|
| 97 |
+
{%- endif %}
|
| 98 |
+
{%- endif %}
|
| 99 |
+
{%- set reasoning_content = reasoning_content|trim %}
|
| 100 |
+
{%- if loop.index0 > ns.last_query_index %}
|
| 101 |
+
{{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content + '\n</think>\n\n' + content }}
|
| 102 |
+
{%- else %}
|
| 103 |
+
{{- '<|im_start|>' + message.role + '\n' + content }}
|
| 104 |
+
{%- endif %}
|
| 105 |
+
{%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
|
| 106 |
+
{%- for tool_call in message.tool_calls %}
|
| 107 |
+
{%- if tool_call.function is defined %}
|
| 108 |
+
{%- set tool_call = tool_call.function %}
|
| 109 |
+
{%- endif %}
|
| 110 |
+
{%- if loop.first %}
|
| 111 |
+
{%- if content|trim %}
|
| 112 |
+
{{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
| 113 |
+
{%- else %}
|
| 114 |
+
{{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
| 115 |
+
{%- endif %}
|
| 116 |
+
{%- else %}
|
| 117 |
+
{{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
|
| 118 |
+
{%- endif %}
|
| 119 |
+
{%- if tool_call.arguments is defined %}
|
| 120 |
+
{%- for args_name, args_value in tool_call.arguments|items %}
|
| 121 |
+
{{- '<parameter=' + args_name + '>\n' }}
|
| 122 |
+
{%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
|
| 123 |
+
{{- args_value }}
|
| 124 |
+
{{- '\n</parameter>\n' }}
|
| 125 |
+
{%- endfor %}
|
| 126 |
+
{%- endif %}
|
| 127 |
+
{{- '</function>\n</tool_call>' }}
|
| 128 |
+
{%- endfor %}
|
| 129 |
+
{%- endif %}
|
| 130 |
+
{{- '<|im_end|>\n' }}
|
| 131 |
+
{%- elif message.role == "tool" %}
|
| 132 |
+
{%- if loop.previtem and loop.previtem.role != "tool" %}
|
| 133 |
+
{{- '<|im_start|>user' }}
|
| 134 |
+
{%- endif %}
|
| 135 |
+
{{- '\n<tool_response>\n' }}
|
| 136 |
+
{{- content }}
|
| 137 |
+
{{- '\n</tool_response>' }}
|
| 138 |
+
{%- if not loop.last and loop.nextitem.role != "tool" %}
|
| 139 |
+
{{- '<|im_end|>\n' }}
|
| 140 |
+
{%- elif loop.last %}
|
| 141 |
+
{{- '<|im_end|>\n' }}
|
| 142 |
+
{%- endif %}
|
| 143 |
+
{%- else %}
|
| 144 |
+
{{- raise_exception('Unexpected message role.') }}
|
| 145 |
+
{%- endif %}
|
| 146 |
+
{%- endfor %}
|
| 147 |
+
{%- if add_generation_prompt %}
|
| 148 |
+
{{- '<|im_start|>assistant\n' }}
|
| 149 |
+
{%- if enable_thinking is defined and enable_thinking is false %}
|
| 150 |
+
{{- '<think>\n\n</think>\n\n' }}
|
| 151 |
+
{%- else %}
|
| 152 |
+
{{- '<think>\n' }}
|
| 153 |
+
{%- endif %}
|
| 154 |
+
{%- endif %}
|
config.json
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"Qwen3_5ForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"attn_output_gate": true,
|
| 8 |
+
"bos_token_id": null,
|
| 9 |
+
"dtype": "bfloat16",
|
| 10 |
+
"eos_token_id": 248044,
|
| 11 |
+
"full_attention_interval": 4,
|
| 12 |
+
"head_dim": 256,
|
| 13 |
+
"hidden_act": "silu",
|
| 14 |
+
"hidden_size": 2560,
|
| 15 |
+
"initializer_range": 0.02,
|
| 16 |
+
"intermediate_size": 9216,
|
| 17 |
+
"layer_types": [
|
| 18 |
+
"linear_attention",
|
| 19 |
+
"linear_attention",
|
| 20 |
+
"linear_attention",
|
| 21 |
+
"full_attention",
|
| 22 |
+
"linear_attention",
|
| 23 |
+
"linear_attention",
|
| 24 |
+
"linear_attention",
|
| 25 |
+
"full_attention",
|
| 26 |
+
"linear_attention",
|
| 27 |
+
"linear_attention",
|
| 28 |
+
"linear_attention",
|
| 29 |
+
"full_attention",
|
| 30 |
+
"linear_attention",
|
| 31 |
+
"linear_attention",
|
| 32 |
+
"linear_attention",
|
| 33 |
+
"full_attention",
|
| 34 |
+
"linear_attention",
|
| 35 |
+
"linear_attention",
|
| 36 |
+
"linear_attention",
|
| 37 |
+
"full_attention",
|
| 38 |
+
"linear_attention",
|
| 39 |
+
"linear_attention",
|
| 40 |
+
"linear_attention",
|
| 41 |
+
"full_attention",
|
| 42 |
+
"linear_attention",
|
| 43 |
+
"linear_attention",
|
| 44 |
+
"linear_attention",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"linear_attention",
|
| 47 |
+
"linear_attention",
|
| 48 |
+
"linear_attention",
|
| 49 |
+
"full_attention"
|
| 50 |
+
],
|
| 51 |
+
"linear_conv_kernel_dim": 4,
|
| 52 |
+
"linear_key_head_dim": 128,
|
| 53 |
+
"linear_num_key_heads": 16,
|
| 54 |
+
"linear_num_value_heads": 32,
|
| 55 |
+
"linear_value_head_dim": 128,
|
| 56 |
+
"mamba_ssm_dtype": "float32",
|
| 57 |
+
"max_position_embeddings": 262144,
|
| 58 |
+
"mlp_only_layers": [],
|
| 59 |
+
"model_type": "qwen3_5_text",
|
| 60 |
+
"mtp_num_hidden_layers": 1,
|
| 61 |
+
"mtp_use_dedicated_embeddings": false,
|
| 62 |
+
"num_attention_heads": 16,
|
| 63 |
+
"num_hidden_layers": 32,
|
| 64 |
+
"num_key_value_heads": 4,
|
| 65 |
+
"pad_token_id": null,
|
| 66 |
+
"partial_rotary_factor": 0.25,
|
| 67 |
+
"rms_norm_eps": 1e-06,
|
| 68 |
+
"rope_parameters": {
|
| 69 |
+
"mrope_interleaved": true,
|
| 70 |
+
"mrope_section": [
|
| 71 |
+
11,
|
| 72 |
+
11,
|
| 73 |
+
10
|
| 74 |
+
],
|
| 75 |
+
"partial_rotary_factor": 0.25,
|
| 76 |
+
"rope_theta": 10000000,
|
| 77 |
+
"rope_type": "default"
|
| 78 |
+
},
|
| 79 |
+
"tie_word_embeddings": true,
|
| 80 |
+
"transformers_version": "5.19.0",
|
| 81 |
+
"use_cache": true,
|
| 82 |
+
"vocab_size": 248320
|
| 83 |
+
}
|
decider/__init__.py
ADDED
|
File without changes
|
decider/__pycache__/__init__.cpython-311.pyc
ADDED
|
Binary file (168 Bytes). View file
|
|
|
decider/__pycache__/infer.cpython-311.pyc
ADDED
|
Binary file (39.4 kB). View file
|
|
|
decider/__pycache__/model.cpython-311.pyc
ADDED
|
Binary file (5.64 kB). View file
|
|
|
decider/__pycache__/prompt.cpython-311.pyc
ADDED
|
Binary file (25.3 kB). View file
|
|
|
decider/__pycache__/systemone.cpython-311.pyc
ADDED
|
Binary file (23.3 kB). View file
|
|
|
decider/__pycache__/temperature.cpython-311.pyc
ADDED
|
Binary file (10.9 kB). View file
|
|
|
decider/batching.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Batch planning for the server's queue: which queued rows share one forward.
|
| 2 |
+
|
| 3 |
+
Pure python, no torch, no engine: `plan_batches` takes the row lengths and two callables that describe the engine's
|
| 4 |
+
bucket grid, and returns the groups to run. `decider.serve.batcher` calls it once per collection.
|
| 5 |
+
|
| 6 |
+
Cost model. A forward of B rows padded to length T costs `overhead + B * T` token-units: `overhead` is the fixed cost of
|
| 7 |
+
one forward (launch, replay, the slot gather and the device-to-host copy) expressed as the number of padded tokens that
|
| 8 |
+
take the same time, and `B * T` is the padded token work. Two rows of different lengths are therefore worth merging
|
| 9 |
+
into the longer row's bucket exactly when the padding they add is cheaper than a second forward's overhead.
|
| 10 |
+
`DEFAULT_MERGE_OVERHEAD_TOKENS` is measured on decider-2b (docs/SERVING.md, section on the batching policy); the server
|
| 11 |
+
reads `DECIDER_MERGE_OVERHEAD_TOKENS` over it.
|
| 12 |
+
|
| 13 |
+
The partition is exact, not greedy. Rows are sorted by padded length descending (stable, so rows of the same bucket keep
|
| 14 |
+
their arrival order) and split into consecutive groups; each group runs at the bucket of its longest member, so a group
|
| 15 |
+
starting at position k costs `overhead + g * T[k]`. A dynamic program over the sorted sequence takes the cheapest split,
|
| 16 |
+
subject to `g <= min(max_batch, max_rows(T))`, with ties going to the larger group. The per-bucket grouping the server
|
| 17 |
+
did before is one of the partitions the program may choose (rows of equal length are adjacent in the sorted order), so
|
| 18 |
+
the planned cost is never above it, and with `overhead = 0` the plan is exactly that grouping.
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
DEFAULT_MERGE_OVERHEAD_TOKENS = 512
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def group_cap(T, max_rows, max_batch):
|
| 25 |
+
"""Rows one forward may take at padded length T: the engine's widest captured batch bucket, and the server's cap."""
|
| 26 |
+
return max(1, min(int(max_batch), int(max_rows(T))))
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def batch_cost(size, T, overhead=DEFAULT_MERGE_OVERHEAD_TOKENS):
|
| 30 |
+
"""Cost of one forward of `size` rows padded to `T`, in token-units."""
|
| 31 |
+
return overhead + size * T
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def plan_cost(groups, overhead=DEFAULT_MERGE_OVERHEAD_TOKENS):
|
| 35 |
+
"""Cost of a plan: the sum over its forwards."""
|
| 36 |
+
return sum(batch_cost(len(idx), T, overhead) for T, idx in groups)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def per_bucket_groups(lengths, pad_len, max_rows, max_batch):
|
| 40 |
+
"""The grouping of 1.1.0: rows of the same padded length, chunked to the cap. The baseline `plan_batches` improves on."""
|
| 41 |
+
return _exact_bucket(list(range(len(lengths))), [pad_len(n) for n in lengths], max_rows, max_batch)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
def _exact_bucket(idx, pads, max_rows, max_batch):
|
| 45 |
+
buckets = {}
|
| 46 |
+
for i in idx:
|
| 47 |
+
buckets.setdefault(pads[i], []).append(i)
|
| 48 |
+
out = []
|
| 49 |
+
for T, members in buckets.items():
|
| 50 |
+
cap = group_cap(T, max_rows, max_batch)
|
| 51 |
+
for k in range(0, len(members), cap):
|
| 52 |
+
out.append((T, members[k:k + cap]))
|
| 53 |
+
return out
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def plan_batches(lengths, pad_len, max_rows, max_batch, overhead=DEFAULT_MERGE_OVERHEAD_TOKENS, mergeable=None):
|
| 57 |
+
"""Partition queued rows into forwards.
|
| 58 |
+
|
| 59 |
+
lengths: token count of every queued row, in arrival order.
|
| 60 |
+
pad_len(n): the engine's padded length for a row of n tokens (`EngineV2.pad_len`).
|
| 61 |
+
max_rows(T):rows the engine runs in one forward at padded length T (`EngineV2.max_rows`).
|
| 62 |
+
max_batch: the server's own cap (`DECIDER_MAX_BATCH`).
|
| 63 |
+
overhead: fixed cost of a forward in token-units (see the module docstring).
|
| 64 |
+
mergeable(n): False for a row that must not be padded into another row's bucket. Rows above the last captured
|
| 65 |
+
length bucket run eager at a request-specific shape, so they are grouped by exact length as before.
|
| 66 |
+
|
| 67 |
+
-> [(padded length, [row indices]), ...]. Every row appears in exactly one group; a group's padded length is at
|
| 68 |
+
least every member's own padded length; a group holds at most `min(max_batch, max_rows(T))` rows.
|
| 69 |
+
"""
|
| 70 |
+
n = len(lengths)
|
| 71 |
+
if n == 0:
|
| 72 |
+
return []
|
| 73 |
+
pads = [pad_len(x) for x in lengths]
|
| 74 |
+
merge = [True] * n if mergeable is None else [bool(mergeable(x)) for x in lengths]
|
| 75 |
+
out = _exact_bucket([i for i in range(n) if not merge[i]], pads, max_rows, max_batch)
|
| 76 |
+
|
| 77 |
+
order = sorted((i for i in range(n) if merge[i]), key=lambda i: -pads[i]) # stable: equal buckets keep arrival order
|
| 78 |
+
m = len(order)
|
| 79 |
+
dp = [0] * (m + 1) # dp[k]: cheapest cost of covering order[k:]
|
| 80 |
+
take = [0] * (m + 1)
|
| 81 |
+
for k in range(m - 1, -1, -1):
|
| 82 |
+
T = pads[order[k]]
|
| 83 |
+
cap = min(group_cap(T, max_rows, max_batch), m - k)
|
| 84 |
+
best, best_g = None, 1
|
| 85 |
+
for g in range(1, cap + 1):
|
| 86 |
+
c = overhead + g * T + dp[k + g]
|
| 87 |
+
if best is None or c <= best: # ties go to the larger group: same modelled cost, one fewer launch
|
| 88 |
+
best, best_g = c, g
|
| 89 |
+
dp[k], take[k] = best, best_g
|
| 90 |
+
k = 0
|
| 91 |
+
while k < m:
|
| 92 |
+
g = take[k]
|
| 93 |
+
out.append((pads[order[k]], sorted(order[k:k + g]))) # inside a group, arrival order
|
| 94 |
+
k += g
|
| 95 |
+
return out
|
decider/calibrate.py
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Fit decider_config.json "temperature_by_type" by NLL, one temperature per answer type (decider.temperature).
|
| 2 |
+
|
| 3 |
+
python -m decider.calibrate records.jsonl [--min-rows 50]
|
| 4 |
+
-> {"temperature_by_type": {"choice": 1.48, "noul": 2.22, "score": 1.38}, "temperature": 1.6, "rows": {...}, "nll": {...}}
|
| 5 |
+
|
| 6 |
+
A record is one answer with its gold, read at temperature 1:
|
| 7 |
+
{"type": "choice" | "noul" | "score", "logits": [one value per option], "gold": option index}
|
| 8 |
+
{"type": "score", "level_logits": [[no, yes], one pair per level], "gold": level index} Score with isolated levels
|
| 9 |
+
"probs" / "level_probs" (probabilities at temperature 1) may stand in for "logits" / "level_logits": log p is the logit up to a
|
| 10 |
+
constant, which the softmax ignores. A noul gold is 1 for true, 0 for false. Every record is checked first (check_record):
|
| 11 |
+
a malformed one (gold out of range or not an integer, wrong shape, NaN, a row without a finite logit, an isolated Score whose
|
| 12 |
+
levels all have P(yes) = 0) raises ValueError naming its position.
|
| 13 |
+
|
| 14 |
+
An isolated-levels Score answer is fitted through the readout the server uses: every level row is softmax(row / T), and the
|
| 15 |
+
answer is P(yes) of each level divided by their sum (decider.systemone.combine_isolated). So the "score" temperature is fitted on
|
| 16 |
+
Score answers whichever readout the model serves them with; fit it on records collected with the same isolated_levels setting.
|
| 17 |
+
|
| 18 |
+
`collect(decider, examples)` produces records from a decider.infer.Decider and labelled /v1/systemone-shaped examples.
|
| 19 |
+
"temperature" in the output is one temperature fitted on all records together, for comparison; a type with fewer than
|
| 20 |
+
--min-rows records is left out of the map (it then uses "temperature" of the config).
|
| 21 |
+
"""
|
| 22 |
+
import json, math, sys
|
| 23 |
+
import numpy as np
|
| 24 |
+
|
| 25 |
+
from decider.temperature import TYPES
|
| 26 |
+
|
| 27 |
+
GRID = np.exp(np.linspace(math.log(0.05), math.log(20.0), 801)) # 0.05 .. 20, about 0.75 % apart
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _log(x):
|
| 31 |
+
x = np.asarray(x, dtype=np.float64)
|
| 32 |
+
with np.errstate(divide="ignore"):
|
| 33 |
+
return np.where(x > 0, np.log(np.clip(x, 1e-300, None)), -np.inf)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _logits(rec, key):
|
| 37 |
+
if key in rec:
|
| 38 |
+
return np.asarray(rec[key], dtype=np.float64)
|
| 39 |
+
alt = {"logits": "probs", "level_logits": "level_probs"}[key]
|
| 40 |
+
return _log(rec[alt])
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def _lse(z, axis=-1):
|
| 44 |
+
m = np.max(z, axis=axis, keepdims=True)
|
| 45 |
+
m = np.where(np.isfinite(m), m, 0.0)
|
| 46 |
+
with np.errstate(divide="ignore"): # an all -inf row gives -inf, handled by the callers
|
| 47 |
+
return (m + np.log(np.sum(np.exp(z - m), axis=axis, keepdims=True))).squeeze(axis)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def nll(rec, T):
|
| 51 |
+
"""Negative log-likelihood of one record's gold at temperature T."""
|
| 52 |
+
return float(_Batch([rec]).nll(T))
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _isolated(rec):
|
| 56 |
+
return "level_logits" in rec or "level_probs" in rec
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def check_record(rec, i=None):
|
| 60 |
+
"""Raise ValueError unless `rec` is a well-formed record (module docstring): an integral gold within the options or levels,
|
| 61 |
+
a 1-D row of at least two options or an [levels, 2] array of at least two levels, logits that are not NaN or +inf,
|
| 62 |
+
probabilities in [0, 1] with positive mass, and isolated levels only on a Score record."""
|
| 63 |
+
where = f"record {i}" if i is not None else "record"
|
| 64 |
+
if not isinstance(rec, dict):
|
| 65 |
+
raise ValueError(f"{where}: not a JSON object")
|
| 66 |
+
if rec.get("type") not in TYPES:
|
| 67 |
+
raise ValueError(f"{where}: type {rec.get('type')!r}, expected one of {', '.join(TYPES)}")
|
| 68 |
+
iso = _isolated(rec)
|
| 69 |
+
keys = ("level_logits", "level_probs") if iso else ("logits", "probs")
|
| 70 |
+
if sum(k in rec for k in keys) != 1 or (iso and ("logits" in rec or "probs" in rec)):
|
| 71 |
+
raise ValueError(f"{where}: give exactly one of logits, probs, level_logits, level_probs")
|
| 72 |
+
if iso and rec["type"] != "score":
|
| 73 |
+
raise ValueError(f"{where}: isolated levels (level_logits / level_probs) belong to a score record, not {rec['type']!r}")
|
| 74 |
+
key = keys[0] if keys[0] in rec else keys[1]
|
| 75 |
+
try:
|
| 76 |
+
a = np.asarray(rec[key], dtype=np.float64)
|
| 77 |
+
except (TypeError, ValueError):
|
| 78 |
+
raise ValueError(f"{where}: {key} is not a numeric array") from None
|
| 79 |
+
if iso and (a.ndim != 2 or a.shape[1] != 2 or a.shape[0] < 2):
|
| 80 |
+
raise ValueError(f"{where}: {key} must be one [no, yes] pair per level, at least two levels; got shape {a.shape}")
|
| 81 |
+
if not iso and (a.ndim != 1 or a.shape[0] < 2):
|
| 82 |
+
raise ValueError(f"{where}: {key} must be one value per option, at least two options; got shape {a.shape}")
|
| 83 |
+
if key.endswith("probs"):
|
| 84 |
+
if not np.all(np.isfinite(a)) or np.any(a < 0) or np.any(a > 1) or np.any(a.sum(-1) <= 0):
|
| 85 |
+
raise ValueError(f"{where}: {key} must be probabilities in [0, 1] with positive mass per row")
|
| 86 |
+
elif np.any(np.isnan(a)) or np.any(a == np.inf) or not np.all(np.any(np.isfinite(a), axis=-1)):
|
| 87 |
+
raise ValueError(f"{where}: {key} contains NaN or +inf, or a row without a finite value")
|
| 88 |
+
if rec["type"] == "noul" and a.shape[0] != 2:
|
| 89 |
+
raise ValueError(f"{where}: a noul record has exactly two options (false, true); got {a.shape[0]}")
|
| 90 |
+
if iso and not np.any(a[:, 1] > 0 if key == "level_probs" else np.isfinite(a[:, 1])):
|
| 91 |
+
raise ValueError(f"{where}: no level has a yes probability above 0, so the served Score answer has no distribution")
|
| 92 |
+
g = rec.get("gold")
|
| 93 |
+
if isinstance(g, bool) or not isinstance(g, (int, np.integer)) or not 0 <= g < a.shape[0]:
|
| 94 |
+
raise ValueError(f"{where}: gold must be an integer index in 0..{a.shape[0] - 1}, got {g!r}")
|
| 95 |
+
return rec
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
class _Batch:
|
| 99 |
+
"""Records packed into padded arrays, so one temperature is scored over all of them at once."""
|
| 100 |
+
def __init__(self, records):
|
| 101 |
+
records = [check_record(r, i) for i, r in enumerate(records)]
|
| 102 |
+
lists = [r for r in records if not _isolated(r)]
|
| 103 |
+
isos = [r for r in records if _isolated(r)]
|
| 104 |
+
self.n = len(lists) + len(isos)
|
| 105 |
+
self.L = self.Lg = self.I = self.Im = self.Ig = None
|
| 106 |
+
if lists:
|
| 107 |
+
zs = [_logits(r, "logits") for r in lists]
|
| 108 |
+
K = max(len(z) for z in zs)
|
| 109 |
+
self.L = np.full((len(zs), K), -np.inf)
|
| 110 |
+
for i, z in enumerate(zs):
|
| 111 |
+
self.L[i, :len(z)] = z
|
| 112 |
+
self.Lg = np.array([int(r["gold"]) for r in lists])
|
| 113 |
+
if isos:
|
| 114 |
+
zs = [_logits(r, "level_logits").reshape(-1, 2) for r in isos]
|
| 115 |
+
n = max(len(z) for z in zs)
|
| 116 |
+
self.I = np.zeros((len(zs), n, 2)); self.Im = np.zeros((len(zs), n), dtype=bool)
|
| 117 |
+
for i, z in enumerate(zs):
|
| 118 |
+
self.I[i, :len(z)] = z; self.Im[i, :len(z)] = True
|
| 119 |
+
self.Ig = np.array([int(r["gold"]) for r in isos])
|
| 120 |
+
|
| 121 |
+
def nll(self, T):
|
| 122 |
+
"""Sum of the records' NLL at temperature T."""
|
| 123 |
+
tot = 0.0
|
| 124 |
+
if self.L is not None:
|
| 125 |
+
z = self.L / T
|
| 126 |
+
zg = z[np.arange(len(z)), self.Lg]
|
| 127 |
+
tot += float(np.sum(np.where(np.isfinite(zg), _lse(z) - zg, 690.0)))
|
| 128 |
+
if self.I is not None: # in log space: P(yes) of very confident rows underflows otherwise
|
| 129 |
+
z = self.I / T
|
| 130 |
+
with np.errstate(invalid="ignore"):
|
| 131 |
+
lp = np.where(self.Im, z[..., 1] - _lse(z), -np.inf) # log P(yes) per level, -inf on padding
|
| 132 |
+
norm = _lse(lp) # log of the summed P(yes)
|
| 133 |
+
lg = lp[np.arange(len(z)), self.Ig]
|
| 134 |
+
with np.errstate(invalid="ignore"):
|
| 135 |
+
per = np.where(np.isfinite(lg), norm - lg, 690.0) # check_record guarantees a finite norm
|
| 136 |
+
tot += float(np.sum(per))
|
| 137 |
+
return tot
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def mean_nll(records, T):
|
| 141 |
+
b = records if isinstance(records, _Batch) else _Batch(records)
|
| 142 |
+
return b.nll(T) / b.n
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
def fit(records, grid=GRID):
|
| 146 |
+
"""The temperature with the lowest mean NLL on the grid, refined by a golden-section search between its neighbours."""
|
| 147 |
+
records = records if isinstance(records, _Batch) else _Batch(list(records))
|
| 148 |
+
if not records.n:
|
| 149 |
+
raise ValueError("no records to fit")
|
| 150 |
+
vals = [mean_nll(records, T) for T in grid]
|
| 151 |
+
i = int(np.argmin(vals))
|
| 152 |
+
lo, hi = math.log(grid[max(i - 1, 0)]), math.log(grid[min(i + 1, len(grid) - 1)])
|
| 153 |
+
r = (math.sqrt(5) - 1) / 2
|
| 154 |
+
a, b = hi - r * (hi - lo), lo + r * (hi - lo)
|
| 155 |
+
fa, fb = mean_nll(records, math.exp(a)), mean_nll(records, math.exp(b))
|
| 156 |
+
for _ in range(40):
|
| 157 |
+
if fa < fb:
|
| 158 |
+
hi, b, fb = b, a, fa; a = hi - r * (hi - lo); fa = mean_nll(records, math.exp(a))
|
| 159 |
+
else:
|
| 160 |
+
lo, a, fa = a, b, fb; b = lo + r * (hi - lo); fb = mean_nll(records, math.exp(b))
|
| 161 |
+
T = math.exp((lo + hi) / 2)
|
| 162 |
+
return T if mean_nll(records, T) <= vals[i] else float(grid[i])
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
def fit_by_type(records, min_rows=50):
|
| 166 |
+
"""-> {"temperature_by_type": {type: T}, "temperature": pooled T, "rows": {type: n}, "nll": {type: {"T=1": .., "fitted": ..}}}"""
|
| 167 |
+
records = [check_record(r, i) for i, r in enumerate(records)]
|
| 168 |
+
groups = {t: [] for t in TYPES}
|
| 169 |
+
for r in records:
|
| 170 |
+
groups[r["type"]].append(r)
|
| 171 |
+
out = {"temperature_by_type": {}, "temperature": round(fit(records), 3) if records else None,
|
| 172 |
+
"rows": {t: len(g) for t, g in groups.items()}, "nll": {}}
|
| 173 |
+
for t, g in groups.items():
|
| 174 |
+
if len(g) < max(1, min_rows):
|
| 175 |
+
continue
|
| 176 |
+
b = _Batch(g); T = fit(b)
|
| 177 |
+
out["temperature_by_type"][t] = round(T, 3)
|
| 178 |
+
out["nll"][t] = {"T=1": round(mean_nll(b, 1.0), 4), "fitted": round(mean_nll(b, T), 4)}
|
| 179 |
+
return out
|
| 180 |
+
|
| 181 |
+
|
| 182 |
+
def collect(decider, examples, independent=True):
|
| 183 |
+
"""Records for fit_by_type from a decider.infer.Decider, read at temperature 1 on the uncached state-first path.
|
| 184 |
+
examples: iterable of (state, questions, golds) with questions as /v1/systemone takes them and golds {question id: gold},
|
| 185 |
+
the gold being a Choice option name, a Noul true/false, or a Score level index. Questions without a gold are skipped.
|
| 186 |
+
Score questions are read with the Decider's isolated_levels setting, as system_one serves them."""
|
| 187 |
+
recs = []
|
| 188 |
+
for state, questions, golds in examples:
|
| 189 |
+
rqs, index, items = decider._system_one_items(state, questions, independent, layout="state_first")
|
| 190 |
+
rows = decider._system_one_probs(items, "state_first", temperature=1.0)
|
| 191 |
+
for k, kind, s, n in index:
|
| 192 |
+
if k not in golds:
|
| 193 |
+
continue
|
| 194 |
+
rq, g = rqs[k], golds[k]
|
| 195 |
+
if rq["type"] == "noul" and not (isinstance(g, bool) or g in (0, 1)):
|
| 196 |
+
raise ValueError(f"gold of noul question {k!r} must be true or false, got {g!r}")
|
| 197 |
+
if rq["type"] == "score" and (isinstance(g, bool) or not isinstance(g, int) or not 0 <= g < len(rq["names"])):
|
| 198 |
+
raise ValueError(f"gold of score question {k!r} must be a level index in 0..{len(rq['names']) - 1}, got {g!r}")
|
| 199 |
+
key = bool(g) if rq["type"] == "noul" else g
|
| 200 |
+
if key not in rq["names"]:
|
| 201 |
+
raise ValueError(f"gold of choice question {k!r} is not one of its options: {g!r}")
|
| 202 |
+
gold = rq["names"].index(key)
|
| 203 |
+
if kind == "iso":
|
| 204 |
+
recs.append({"type": rq["type"], "level_probs": [rows[s + j][:2] for j in range(n)], "gold": gold})
|
| 205 |
+
else:
|
| 206 |
+
recs.append({"type": rq["type"], "probs": rows[s][:len(rq["options"])], "gold": gold})
|
| 207 |
+
return recs
|
| 208 |
+
|
| 209 |
+
|
| 210 |
+
def main(argv=None):
|
| 211 |
+
import argparse
|
| 212 |
+
ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
|
| 213 |
+
ap.add_argument("records", help="JSON lines, one record per line (see the module docstring)")
|
| 214 |
+
ap.add_argument("--min-rows", type=int, default=50, help="fit a type only when it has at least this many records")
|
| 215 |
+
a = ap.parse_args(argv)
|
| 216 |
+
with open(a.records) as f:
|
| 217 |
+
recs = [json.loads(line) for line in f if line.strip()]
|
| 218 |
+
print(json.dumps(fit_by_type(recs, a.min_rows), indent=1))
|
| 219 |
+
|
| 220 |
+
|
| 221 |
+
if __name__ == "__main__":
|
| 222 |
+
main(sys.argv[1:])
|
decider/engine.py
ADDED
|
@@ -0,0 +1,216 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Low-latency inference engine: shape-bucketed CUDA graphs over the one-pass decision model.
|
| 2 |
+
|
| 3 |
+
Right padding + causal layers => pad positions never influence earlier slots, so no attention
|
| 4 |
+
mask is needed and every (B, T) bucket can be captured once and replayed. The graph outputs
|
| 5 |
+
option-letter logits for all positions [B, T, K]; slots are gathered outside.
|
| 6 |
+
"""
|
| 7 |
+
import time, torch, torch._dynamo, torch.nn.functional as F
|
| 8 |
+
from decider.model import DecisionModel, collate
|
| 9 |
+
from decider.prompt import build, MAX_OPTIONS
|
| 10 |
+
from decider.temperature import scaled_softmax, slot_temperatures
|
| 11 |
+
|
| 12 |
+
T_BUCKETS = [64, 128, 192, 256, 320, 384, 512, 640, 768, 1024, 1280, 1536, 2048]
|
| 13 |
+
B_BUCKETS = [1, 2, 4, 8, 16, 32, 64]
|
| 14 |
+
GRAPH_MAX_T = 2048 # longer inputs (up to the 32k request budget) run eagerly: compute dominates there, and one graph
|
| 15 |
+
LONG_STEP = 1024 # per (B, T) shape would cost a compile + capture for every new length
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def _bucket(x, buckets):
|
| 19 |
+
for b in buckets:
|
| 20 |
+
if x <= b:
|
| 21 |
+
return b
|
| 22 |
+
return None
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def fused_causal_conv1d_fn(hidden_states, weight, bias=None, activation=None, **kwargs):
|
| 26 |
+
"""Depthwise causal conv (kernel k) as k shifted multiply-adds: fuses under torch.compile,
|
| 27 |
+
unlike the cuDNN grouped conv fallback (which was ~11% of batched GPU time)."""
|
| 28 |
+
B, C, T = hidden_states.shape; k = weight.shape[-1]
|
| 29 |
+
x = F.pad(hidden_states.to(weight.dtype), (k - 1, 0))
|
| 30 |
+
out = x[:, :, k - 1:k - 1 + T] * weight[:, k - 1][None, :, None]
|
| 31 |
+
for j in range(k - 1):
|
| 32 |
+
out = out + x[:, :, j:j + T] * weight[:, j][None, :, None]
|
| 33 |
+
if bias is not None:
|
| 34 |
+
out = out + bias[None, :, None]
|
| 35 |
+
if activation == "silu":
|
| 36 |
+
out = F.silu(out)
|
| 37 |
+
elif activation is not None:
|
| 38 |
+
from transformers.activations import ACT2FN
|
| 39 |
+
out = ACT2FN[activation](out)
|
| 40 |
+
return out.to(hidden_states.dtype)
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def patch_conv():
|
| 44 |
+
from transformers.models.qwen3_5 import modeling_qwen3_5 as mq
|
| 45 |
+
mq.causal_conv1d_fn = fused_causal_conv1d_fn
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def read_slots(out, rows, slots, nopts, temperature, n_per_item):
|
| 49 |
+
"""One gather + one softmax + one device-to-host copy for the whole batch (was: three small kernels and a sync per item).
|
| 50 |
+
out [B, T, K] logits; rows/slots/nopts: flat python lists, one entry per question; n_per_item: questions per item.
|
| 51 |
+
temperature: a number for every question, or a flat list with one temperature per question (decider.temperature)."""
|
| 52 |
+
dev = out.device; idx = torch.tensor([rows, slots, nopts], dtype=torch.long).to(dev, non_blocking=True)
|
| 53 |
+
lg = out[idx[0], idx[1]] # [N, K]
|
| 54 |
+
lg = lg.masked_fill(torch.arange(lg.shape[1], device=dev)[None, :] >= idx[2][:, None], float("-inf"))
|
| 55 |
+
p = scaled_softmax(lg, temperature).cpu()
|
| 56 |
+
return list(torch.split(p, n_per_item))
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def fill_ids(items_ids, B, T, pad):
|
| 60 |
+
import numpy as np
|
| 61 |
+
a = np.full((B, T), pad, dtype=np.int64)
|
| 62 |
+
for b, x in enumerate(items_ids): a[b, :len(x)] = x
|
| 63 |
+
return torch.from_numpy(a)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def set_attention_backend_policy():
|
| 67 |
+
"""Turn off the cuDNN scaled-dot-product-attention backend. On Blackwell with torch 2.14 / CUDA 13 it returns wrong,
|
| 68 |
+
finite output for masked rectangular attention, which is what the shared-state path (`Engine.score_shared`) and the schema
|
| 69 |
+
cache run when a suffix is scored against a cached prefix; the math and memory-efficient backends are correct. Measured on
|
| 70 |
+
a Decision Index row: the cached path answered a wrong option at p=0.93 where the full forward and the corrected cached path
|
| 71 |
+
both give option_4 at p=0.95 (decider2/SERVING_V2_REVIEW.md in the research notes). Must run before torch.compile and CUDA
|
| 72 |
+
graph capture: captured graphs keep the backend they were captured with."""
|
| 73 |
+
if hasattr(torch.backends.cuda, "enable_cudnn_sdp"):
|
| 74 |
+
torch.backends.cuda.enable_cudnn_sdp(False)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class Engine:
|
| 78 |
+
"""compile: torch.compile the forward (needs use_cache=False; ~1.4x batched, fuses elementwise work).
|
| 79 |
+
fp8: e4m3 weights + per-token activation scaling on the big linears (Hopper tensor cores).
|
| 80 |
+
conv_patch: fusable depthwise causal conv instead of the cuDNN fallback."""
|
| 81 |
+
def __init__(self, path, device="cuda", dtype=torch.bfloat16, use_graphs=True, max_ctx_tokens=1536,
|
| 82 |
+
compile=True, fp8=False, conv_patch=True):
|
| 83 |
+
set_attention_backend_policy()
|
| 84 |
+
if conv_patch:
|
| 85 |
+
if str(device).startswith("mps"):
|
| 86 |
+
from decider.mps_ops import patch_mps
|
| 87 |
+
patch_mps()
|
| 88 |
+
else:
|
| 89 |
+
patch_conv()
|
| 90 |
+
self.m = DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
|
| 91 |
+
use_graphs = use_graphs and torch.device(device).type == "cuda"
|
| 92 |
+
self.tok = self.m.tok; self.dev = device; self.use_graphs = use_graphs; self.max_ctx = max_ctx_tokens
|
| 93 |
+
self.core, self.W = self.m.lm.model, self.m.lm.lm_head.weight[self.m.letters].detach().clone()
|
| 94 |
+
self.cfg = dict(compile=compile, fp8=fp8, conv_patch=conv_patch, graphs=use_graphs)
|
| 95 |
+
if fp8:
|
| 96 |
+
from decider.fp8 import convert_to_fp8
|
| 97 |
+
self.cfg["fp8_layers"] = convert_to_fp8(self.core)
|
| 98 |
+
if compile:
|
| 99 |
+
torch._dynamo.config.cache_size_limit = 128
|
| 100 |
+
self._fwd_impl = torch.compile(self._fwd_eager, dynamic=False)
|
| 101 |
+
else:
|
| 102 |
+
self._fwd_impl = self._fwd_eager
|
| 103 |
+
self.graphs = {} # (B, T) -> (static_ids, static_out, graph)
|
| 104 |
+
self.pool = torch.cuda.graph_pool_handle() if (use_graphs and str(device).startswith("cuda")) else None
|
| 105 |
+
self.stats = dict(graph_captures=0, forwards=0)
|
| 106 |
+
|
| 107 |
+
def _fwd_eager(self, ids):
|
| 108 |
+
h = self.core(input_ids=ids, use_cache=False).last_hidden_state
|
| 109 |
+
return F.linear(h, self.W).float() # [B, T, K]
|
| 110 |
+
|
| 111 |
+
@torch.no_grad()
|
| 112 |
+
def _fwd(self, ids):
|
| 113 |
+
return self._fwd_impl(ids)
|
| 114 |
+
|
| 115 |
+
def _capture(self, B, T):
|
| 116 |
+
s_ids = torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev)
|
| 117 |
+
st = torch.cuda.Stream(); st.wait_stream(torch.cuda.current_stream())
|
| 118 |
+
with torch.cuda.stream(st):
|
| 119 |
+
for _ in range(3): self._fwd(s_ids) # warm-up: compile / triton autotune
|
| 120 |
+
torch.cuda.current_stream().wait_stream(st)
|
| 121 |
+
g = torch.cuda.CUDAGraph()
|
| 122 |
+
with torch.cuda.graph(g, pool=self.pool):
|
| 123 |
+
s_out = self._fwd(s_ids)
|
| 124 |
+
self.stats["graph_captures"] += 1
|
| 125 |
+
return s_ids, s_out, g
|
| 126 |
+
|
| 127 |
+
@torch.no_grad()
|
| 128 |
+
def logits_all(self, ids):
|
| 129 |
+
"""ids: [B, T] long on device (already right-padded to a bucket). Returns [B, T, K] float."""
|
| 130 |
+
B, T = ids.shape; self.stats["forwards"] += 1
|
| 131 |
+
if T > GRAPH_MAX_T:
|
| 132 |
+
self.stats["long_forwards"] = self.stats.get("long_forwards", 0) + 1
|
| 133 |
+
return self._fwd_eager(ids)
|
| 134 |
+
if not self.use_graphs:
|
| 135 |
+
return self._fwd(ids)
|
| 136 |
+
key = (B, T)
|
| 137 |
+
if key not in self.graphs:
|
| 138 |
+
self.graphs[key] = self._capture(B, T)
|
| 139 |
+
s_ids, s_out, g = self.graphs[key]
|
| 140 |
+
s_ids.copy_(ids); g.replay()
|
| 141 |
+
return s_out
|
| 142 |
+
|
| 143 |
+
@torch.no_grad()
|
| 144 |
+
def score_items(self, items, temperature=1.0):
|
| 145 |
+
"""items: list of dicts from prompt.build. Returns list of [n_q, MAX_OPTIONS] prob tensors (cpu).
|
| 146 |
+
temperature: a number, or one entry per item (a number or one number per slot; decider.temperature.for_items)."""
|
| 147 |
+
Tmax = max(len(it["ids"]) for it in items)
|
| 148 |
+
T = _bucket(Tmax, T_BUCKETS) or -(-Tmax // LONG_STEP) * LONG_STEP
|
| 149 |
+
B = (_bucket(len(items), B_BUCKETS) or len(items)) if T <= GRAPH_MAX_T else len(items)
|
| 150 |
+
ids = fill_ids([it["ids"] for it in items], B, T, self.tok.pad_token_id)
|
| 151 |
+
out = self.logits_all(ids.to(self.dev, non_blocking=True))
|
| 152 |
+
return read_slots(out, [b for b, it in enumerate(items) for _ in it["slots"]], [s for it in items for s in it["slots"]],
|
| 153 |
+
[n for it in items for n in it["nopts"]], slot_temperatures(temperature, items), [len(it["slots"]) for it in items])
|
| 154 |
+
|
| 155 |
+
@torch.no_grad()
|
| 156 |
+
def score_shared(self, items, temperature=1.0, min_prefix=192):
|
| 157 |
+
"""Rows that start with the same tokens (one state, one question per row): run the shared prefix once, fork its
|
| 158 |
+
cache (attention KV + delta-net conv/recurrent states), and run only the question suffixes.
|
| 159 |
+
Same answers as score_items up to kernel round-off; cost ~ state + sum(questions) instead of n * state.
|
| 160 |
+
The fork is made in chunks that fit `DECIDER_SHARED_FORK_GB`, so the peak memory does not grow with the question
|
| 161 |
+
count; the implementation is decider.shared_prefix, shared with EngineV2."""
|
| 162 |
+
from decider import shared_prefix # imported here: decider.shared_prefix imports this module
|
| 163 |
+
out = shared_prefix.score_shared(self, items, temperature, min_prefix)
|
| 164 |
+
if out is None:
|
| 165 |
+
return self.score_items(items, temperature)
|
| 166 |
+
self.stats["shared_prefix_calls"] = self.stats.get("shared_prefix_calls", 0) + 1
|
| 167 |
+
return out
|
| 168 |
+
|
| 169 |
+
def warmup(self, shapes=((1, 128), (1, 256), (1, 384), (1, 512), (8, 256), (8, 512), (32, 256), (32, 512))):
|
| 170 |
+
t = time.time()
|
| 171 |
+
for B, T in shapes:
|
| 172 |
+
self.logits_all(torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev))
|
| 173 |
+
torch.cuda.synchronize(); return time.time() - t
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
if __name__ == "__main__":
|
| 177 |
+
import sys, random, numpy as np
|
| 178 |
+
from decider import data as D
|
| 179 |
+
from decider.infer import Decider
|
| 180 |
+
path = sys.argv[1] if len(sys.argv) > 1 else "runs/r3_v2/model"
|
| 181 |
+
cfg = dict(compile="nocompile" not in sys.argv[2:], fp8="fp8" in sys.argv[2:], conv_patch="noconv" not in sys.argv[2:])
|
| 182 |
+
_, evals = D.load_cache("data/tasks.pkl")
|
| 183 |
+
eng = Engine(path, **cfg); print("engine cfg", eng.cfg)
|
| 184 |
+
rng = random.Random(0)
|
| 185 |
+
exs = evals["support_tickets"][:64] + evals["clinc_oos"][:64] + evals["race"][:32]
|
| 186 |
+
from decider.prompt import chat_for_model
|
| 187 |
+
chat = chat_for_model(path, eng.tok)
|
| 188 |
+
items = [build(e, eng.tok, rng, max_ctx_tokens=1536, chat=chat) for e in exs]
|
| 189 |
+
# correctness vs eager masked forward (DecisionModel.slot_logits)
|
| 190 |
+
ref = []
|
| 191 |
+
with torch.no_grad():
|
| 192 |
+
for i in range(0, len(items), 16):
|
| 193 |
+
b = collate(items[i:i + 16], eng.tok.pad_token_id)
|
| 194 |
+
lg = eng.m.slot_logits(b["input_ids"].cuda(), b["attention_mask"].cuda(), b["slot_idx"].cuda(), b["slot_batch"].cuda(), b["nopts"].cuda())
|
| 195 |
+
ref.append(torch.softmax(lg, -1).cpu())
|
| 196 |
+
ref = torch.cat(ref)
|
| 197 |
+
got = torch.cat(eng.score_items(items))
|
| 198 |
+
print(f"max |p_graph - p_eager| = {(ref - got).abs().max():.4f} over {len(ref)} questions; argmax agreement {(ref.argmax(1) == got.argmax(1)).float().mean():.4f}")
|
| 199 |
+
print(f"warmup capture of 8 buckets: {eng.warmup():.1f}s; captures so far {eng.stats['graph_captures']}")
|
| 200 |
+
# latency: single real requests
|
| 201 |
+
for name, pool in [("support_tickets", exs[:64]), ("clinc_oos", exs[64:128]), ("race", exs[128:])]:
|
| 202 |
+
its = [build(e, eng.tok, rng, chat=chat) for e in pool]
|
| 203 |
+
ts = []
|
| 204 |
+
for it in its[:40]:
|
| 205 |
+
torch.cuda.synchronize(); t = time.time(); eng.score_items([it]); torch.cuda.synchronize(); ts.append(time.time() - t)
|
| 206 |
+
ts = np.array(ts[5:]) * 1000
|
| 207 |
+
print(f"single request {name:16s}: p50 {np.median(ts):5.1f} ms p90 {np.percentile(ts, 90):5.1f} ms (avg {np.mean([len(i['ids']) for i in its]):.0f} tok, {len(its[0]['slots'])} q)")
|
| 208 |
+
for bs in (8, 32):
|
| 209 |
+
ts = []
|
| 210 |
+
for i in range(0, min(len(its), bs * 6), bs):
|
| 211 |
+
chunk = its[i:i + bs]
|
| 212 |
+
if len(chunk) < bs: break
|
| 213 |
+
torch.cuda.synchronize(); t = time.time(); eng.score_items(chunk); torch.cuda.synchronize(); ts.append(time.time() - t)
|
| 214 |
+
ts = np.array(ts[1:]) * 1000
|
| 215 |
+
print(f" batch {bs:2d}: p50 {np.median(ts):6.1f} ms -> {bs/np.median(ts)*1000:6.0f} ctx/s, {bs*len(its[0]['slots'])/np.median(ts)*1000:6.0f} decisions/s")
|
| 216 |
+
print("stats", eng.stats, "graphs", len(eng.graphs), f"mem {torch.cuda.memory_reserved()/1e9:.1f} GB")
|
decider/engine_v2.py
ADDED
|
@@ -0,0 +1,213 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Serving engine: every CUDA graph is keyed on (batch bucket, padded length bucket) only.
|
| 2 |
+
|
| 3 |
+
Difference from decider.engine.Engine (the library engine behind decider.infer.Decider):
|
| 4 |
+
* the full bucket grid is captured at start-up by `warmup()`, then `seal()` forbids any further capture, so no request
|
| 5 |
+
can ever pay a graph capture or a torch.compile;
|
| 6 |
+
* the length ladder reaches 8192 instead of 2048, so long rows replay a graph instead of running eager;
|
| 7 |
+
* a batch larger than the widest captured batch bucket for its length is split into captured chunks rather than
|
| 8 |
+
capturing a new shape; rows longer than the last bucket run eager in chunks of at most `token_budget` padded tokens;
|
| 9 |
+
* torch.compile and FP8 are off by default.
|
| 10 |
+
|
| 11 |
+
Numerics are the same as Engine: right padding + causal layers means padded positions never influence earlier slots, so
|
| 12 |
+
the answer for a row does not depend on which (B, T) bucket it was padded into, up to kernel reduction order.
|
| 13 |
+
|
| 14 |
+
The attention backend policy of decider.engine (cuDNN SDPA off) is applied in __init__, before compile and capture. It
|
| 15 |
+
is what makes `score_shared` (a suffix scored against a cached prefix) return the same answers as the full forward; see
|
| 16 |
+
docs/CHANGELOG.md 1.0.2 and tests/test_engine_v2_cuda.py.
|
| 17 |
+
"""
|
| 18 |
+
import time, torch, torch.nn.functional as F
|
| 19 |
+
from decider import shared_prefix
|
| 20 |
+
from decider.engine import read_slots, fill_ids, patch_conv, set_attention_backend_policy
|
| 21 |
+
from decider.model import DecisionModel
|
| 22 |
+
from decider.temperature import item_slice, slot_temperatures
|
| 23 |
+
|
| 24 |
+
T_BUCKETS = [64, 128, 192, 256, 320, 384, 512, 640, 768, 1024, 1280, 1536, 2048, 3072, 4096, 6144, 8192]
|
| 25 |
+
B_BUCKETS = [1, 2, 4, 8, 16, 32]
|
| 26 |
+
TOKEN_BUDGET = 32768 # capture (B, T) only when B * T fits this; B = 1 is always captured
|
| 27 |
+
LONG_STEP = 1024 # rows longer than the last bucket: pad to a multiple of this and run eager
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
class EngineV2:
|
| 31 |
+
"""compile / fp8 default to False. `model` lets a test inject a stand-in DecisionModel instead of loading one."""
|
| 32 |
+
|
| 33 |
+
def __init__(self, path=None, device="cuda", dtype=torch.bfloat16, use_graphs=None, compile=False, fp8=False,
|
| 34 |
+
conv_patch=None, t_buckets=None, b_buckets=None, token_budget=TOKEN_BUDGET, max_ctx_tokens=32768,
|
| 35 |
+
model=None):
|
| 36 |
+
set_attention_backend_policy() # before compile and capture: captured graphs keep their backend
|
| 37 |
+
if conv_patch is None:
|
| 38 |
+
conv_patch = bool(compile) # the unrolled depthwise conv only pays off inside a compiled region
|
| 39 |
+
if str(device).startswith("mps"):
|
| 40 |
+
from decider.mps_ops import patch_mps # the Apple Silicon path of decider.infer.Decider
|
| 41 |
+
patch_mps()
|
| 42 |
+
elif conv_patch:
|
| 43 |
+
patch_conv()
|
| 44 |
+
self.m = model if model is not None else DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
|
| 45 |
+
self.tok = self.m.tok; self.dev = device; self.max_ctx = max_ctx_tokens
|
| 46 |
+
self.core = self.m.lm.model
|
| 47 |
+
self.W = self.m.lm.lm_head.weight[self.m.letters].detach().clone()
|
| 48 |
+
self.t_buckets = sorted(t_buckets or T_BUCKETS)
|
| 49 |
+
self.b_buckets = sorted(b_buckets or B_BUCKETS)
|
| 50 |
+
self.token_budget = int(token_budget)
|
| 51 |
+
self.use_graphs = (str(device).startswith("cuda") if use_graphs is None else bool(use_graphs))
|
| 52 |
+
self.cfg = dict(compile=compile, fp8=fp8, conv_patch=conv_patch, graphs=self.use_graphs,
|
| 53 |
+
t_buckets=self.t_buckets, b_buckets=self.b_buckets, token_budget=self.token_budget)
|
| 54 |
+
if fp8:
|
| 55 |
+
from decider.fp8 import convert_to_fp8
|
| 56 |
+
self.cfg["fp8_layers"] = convert_to_fp8(self.core)
|
| 57 |
+
if compile:
|
| 58 |
+
from torch import _dynamo # not `import torch._dynamo`: that rebinds `torch` locally
|
| 59 |
+
n = len(self.graph_shapes())
|
| 60 |
+
# one specialisation per captured shape, times the frames inside the model. Without the accumulated
|
| 61 |
+
# limits dynamo stops compiling part way through warm-up and silently leaves the rest eager, which with
|
| 62 |
+
# conv_patch on is slower than not compiling at all.
|
| 63 |
+
_dynamo.config.cache_size_limit = max(256, n + 16)
|
| 64 |
+
_dynamo.config.accumulated_cache_size_limit = max(1 << 16, 64 * n)
|
| 65 |
+
for k in ("accumulated_recompile_limit", "recompile_limit"):
|
| 66 |
+
if hasattr(_dynamo.config, k):
|
| 67 |
+
setattr(_dynamo.config, k, max(1 << 16, 64 * n))
|
| 68 |
+
self._fwd_impl = torch.compile(self._fwd_eager, dynamic=False)
|
| 69 |
+
else:
|
| 70 |
+
self._fwd_impl = self._fwd_eager
|
| 71 |
+
self.graphs = {} # (B, T) -> (static ids, static out, graph)
|
| 72 |
+
self.b_for_t = {} # T -> sorted captured batch buckets
|
| 73 |
+
self.pool = torch.cuda.graph_pool_handle() if (self.use_graphs and str(device).startswith("cuda")) else None
|
| 74 |
+
self.sealed = False
|
| 75 |
+
self.stats = dict(graph_captures=0, forwards=0, replays=0, eager_forwards=0, eager_rows=0,
|
| 76 |
+
shared_calls=0, unbucketed_requests=0)
|
| 77 |
+
|
| 78 |
+
# ---- bucket arithmetic -------------------------------------------------
|
| 79 |
+
def graph_shapes(self):
|
| 80 |
+
"""The full grid captured at start-up: every (B, T) whose padded token count fits the budget, plus all of B = 1."""
|
| 81 |
+
return [(B, T) for T in self.t_buckets for B in self.b_buckets if B == 1 or B * T <= self.token_budget]
|
| 82 |
+
|
| 83 |
+
def t_bucket(self, n):
|
| 84 |
+
"""Padded length bucket for a row of n tokens, or None when it is longer than the last bucket."""
|
| 85 |
+
for t in self.t_buckets:
|
| 86 |
+
if n <= t:
|
| 87 |
+
return t
|
| 88 |
+
return None
|
| 89 |
+
|
| 90 |
+
def pad_len(self, n):
|
| 91 |
+
return self.t_bucket(n) or -(-n // LONG_STEP) * LONG_STEP
|
| 92 |
+
|
| 93 |
+
def max_rows(self, T):
|
| 94 |
+
"""Rows per forward at padded length T: the widest captured batch bucket, or, when nothing is captured at T, as
|
| 95 |
+
many rows as fit the token budget (at least 1)."""
|
| 96 |
+
av = self.b_for_t.get(T)
|
| 97 |
+
return av[-1] if av else max(1, self.token_budget // max(T, 1))
|
| 98 |
+
|
| 99 |
+
def b_plan(self, n, T):
|
| 100 |
+
"""Batch sizes covering n rows at length T, each one a captured bucket. One chunk when a bucket fits."""
|
| 101 |
+
av = self.b_for_t.get(T)
|
| 102 |
+
if not av:
|
| 103 |
+
return [n]
|
| 104 |
+
fit = next((b for b in av if b >= n), None)
|
| 105 |
+
if fit is not None:
|
| 106 |
+
return [fit]
|
| 107 |
+
out = []; big = av[-1]
|
| 108 |
+
while n > big:
|
| 109 |
+
out.append(big); n -= big
|
| 110 |
+
out.append(next(b for b in av if b >= n))
|
| 111 |
+
return out
|
| 112 |
+
|
| 113 |
+
# ---- forward -----------------------------------------------------------
|
| 114 |
+
def _fwd_eager(self, ids):
|
| 115 |
+
h = self.core(input_ids=ids, use_cache=False).last_hidden_state
|
| 116 |
+
return F.linear(h, self.W).float() # [B, T, K]
|
| 117 |
+
|
| 118 |
+
@torch.no_grad()
|
| 119 |
+
def _fwd(self, ids):
|
| 120 |
+
return self._fwd_impl(ids)
|
| 121 |
+
|
| 122 |
+
def _capture(self, B, T):
|
| 123 |
+
s_ids = torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev)
|
| 124 |
+
st = torch.cuda.Stream(); st.wait_stream(torch.cuda.current_stream())
|
| 125 |
+
with torch.cuda.stream(st):
|
| 126 |
+
for _ in range(3): self._fwd(s_ids)
|
| 127 |
+
torch.cuda.current_stream().wait_stream(st)
|
| 128 |
+
g = torch.cuda.CUDAGraph()
|
| 129 |
+
with torch.cuda.graph(g, pool=self.pool):
|
| 130 |
+
s_out = self._fwd(s_ids)
|
| 131 |
+
self.stats["graph_captures"] += 1
|
| 132 |
+
return s_ids, s_out, g
|
| 133 |
+
|
| 134 |
+
@torch.no_grad()
|
| 135 |
+
def logits_all(self, ids):
|
| 136 |
+
"""ids: [B, T] long on device, already right-padded to a captured shape where one exists."""
|
| 137 |
+
B, T = ids.shape; self.stats["forwards"] += 1
|
| 138 |
+
if self.use_graphs:
|
| 139 |
+
g = self.graphs.get((B, T))
|
| 140 |
+
if g is None and not self.sealed: # only reachable before seal(); a sealed engine never captures
|
| 141 |
+
g = self.graphs[(B, T)] = self._capture(B, T)
|
| 142 |
+
self.b_for_t.setdefault(T, [])
|
| 143 |
+
if B not in self.b_for_t[T]: self.b_for_t[T] = sorted(self.b_for_t[T] + [B])
|
| 144 |
+
if g is not None:
|
| 145 |
+
s_ids, s_out, gr = g
|
| 146 |
+
s_ids.copy_(ids); gr.replay(); self.stats["replays"] += 1
|
| 147 |
+
return s_out
|
| 148 |
+
self.stats["eager_forwards"] += 1; self.stats["eager_rows"] += B
|
| 149 |
+
return self._fwd_eager(ids) # never the compiled callable: a new shape must not compile
|
| 150 |
+
|
| 151 |
+
# ---- scoring -----------------------------------------------------------
|
| 152 |
+
@torch.no_grad()
|
| 153 |
+
def score_items(self, items, temperature=1.0):
|
| 154 |
+
"""items: dicts from prompt.build / build_rows. -> one [n_q, MAX_OPTIONS] cpu probability tensor per item.
|
| 155 |
+
temperature: a number, or one entry per item (a number or one number per slot; decider.temperature.for_items)."""
|
| 156 |
+
if not items:
|
| 157 |
+
return []
|
| 158 |
+
slot_temperatures(temperature, items) # a length mismatch fails before any forward
|
| 159 |
+
Tmax = max(len(it["ids"]) for it in items)
|
| 160 |
+
T = self.t_bucket(Tmax)
|
| 161 |
+
if T is None:
|
| 162 |
+
self.stats["unbucketed_requests"] += 1
|
| 163 |
+
T = -(-Tmax // LONG_STEP) * LONG_STEP
|
| 164 |
+
per = self.max_rows(T); plan = [per] * (len(items) // per) + ([len(items) % per] if len(items) % per else [])
|
| 165 |
+
else:
|
| 166 |
+
plan = self.b_plan(len(items), T)
|
| 167 |
+
out = []; i = 0
|
| 168 |
+
for B in plan:
|
| 169 |
+
chunk = items[i:i + B]; i += len(chunk)
|
| 170 |
+
ids = fill_ids([it["ids"] for it in chunk], B, T, self.tok.pad_token_id)
|
| 171 |
+
lg = self.logits_all(ids.to(self.dev, non_blocking=True))
|
| 172 |
+
out += read_slots(lg, [b for b, it in enumerate(chunk) for _ in it["slots"]],
|
| 173 |
+
[s for it in chunk for s in it["slots"]], [n for it in chunk for n in it["nopts"]],
|
| 174 |
+
slot_temperatures(item_slice(temperature, i - len(chunk), i), chunk), [len(it["slots"]) for it in chunk])
|
| 175 |
+
return out
|
| 176 |
+
|
| 177 |
+
@torch.no_grad()
|
| 178 |
+
def score_shared(self, items, temperature=1.0, min_prefix=192, budget_bytes=None, rows_per_fork=None):
|
| 179 |
+
"""Rows that start with the same tokens (one state, one question per row): run the shared prefix once, fork its
|
| 180 |
+
cache in chunks that fit a byte budget, run only the question suffixes. The algorithm of Engine.score_shared,
|
| 181 |
+
shared with it in decider.shared_prefix: it is per request, never per schema, so it adds no state that outlives
|
| 182 |
+
the request and no shape that depends on the question set. The prefix and suffix forwards are eager
|
| 183 |
+
(request-specific shapes). Correct only with the cuDNN SDPA backend off, which __init__ arranges."""
|
| 184 |
+
out = shared_prefix.score_shared(self, items, temperature, min_prefix, budget_bytes, rows_per_fork)
|
| 185 |
+
if out is None:
|
| 186 |
+
return self.score_items(items, temperature)
|
| 187 |
+
self.stats["shared_calls"] += 1
|
| 188 |
+
return out
|
| 189 |
+
|
| 190 |
+
# ---- start-up ----------------------------------------------------------
|
| 191 |
+
def warmup(self, shapes=None, log=None):
|
| 192 |
+
"""Capture the whole grid. Call seal() afterwards: from then on an unknown shape runs eager, never captures."""
|
| 193 |
+
t = time.time(); shapes = list(shapes if shapes is not None else self.graph_shapes())
|
| 194 |
+
for j, (B, T) in enumerate(shapes):
|
| 195 |
+
self.logits_all(torch.full((B, T), self.tok.pad_token_id, dtype=torch.long, device=self.dev))
|
| 196 |
+
if log and (j + 1) % 10 == 0:
|
| 197 |
+
log(f"[engine_v2] {j + 1}/{len(shapes)} graphs, {time.time() - t:.0f}s")
|
| 198 |
+
if self.use_graphs:
|
| 199 |
+
torch.cuda.synchronize()
|
| 200 |
+
return time.time() - t
|
| 201 |
+
|
| 202 |
+
def seal(self):
|
| 203 |
+
self.sealed = True
|
| 204 |
+
self.b_for_t = {}
|
| 205 |
+
for B, T in self.graphs:
|
| 206 |
+
self.b_for_t.setdefault(T, []).append(B)
|
| 207 |
+
for T in self.b_for_t:
|
| 208 |
+
self.b_for_t[T].sort()
|
| 209 |
+
return self
|
| 210 |
+
|
| 211 |
+
def describe(self):
|
| 212 |
+
return dict(self.cfg, graphs=len(self.graphs), sealed=self.sealed,
|
| 213 |
+
grid={str(T): self.b_for_t.get(T, []) for T in self.t_buckets})
|
decider/fp8.py
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""FP8 (e4m3) linear layers for Hopper via torch._scaled_mm.
|
| 2 |
+
Weights: per-output-channel scales, quantised once. Activations: per-token dynamic scales.
|
| 3 |
+
Under torch.compile the quantisation ops fuse into the surrounding elementwise work."""
|
| 4 |
+
import torch, torch.nn as nn
|
| 5 |
+
|
| 6 |
+
E4M3_MAX = 448.0
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
def _quant_rowwise(x):
|
| 10 |
+
s = x.abs().amax(dim=-1, keepdim=True).float().clamp(min=1e-12) / E4M3_MAX
|
| 11 |
+
return (x.float() / s).clamp(-E4M3_MAX, E4M3_MAX).to(torch.float8_e4m3fn), s
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class FP8Linear(nn.Module):
|
| 15 |
+
def __init__(self, lin: nn.Linear):
|
| 16 |
+
super().__init__()
|
| 17 |
+
wq, sw = _quant_rowwise(lin.weight.detach()) # [N,K] fp8, [N,1]
|
| 18 |
+
self.register_buffer("wq", wq.contiguous()) # [N,K]; passed as wq.t() -> [K,N] column-major, as _scaled_mm wants
|
| 19 |
+
self.register_buffer("sw_t", sw.t().contiguous()) # [1,N]
|
| 20 |
+
self.bias = None if lin.bias is None else nn.Parameter(lin.bias.detach().clone(), requires_grad=False)
|
| 21 |
+
self.in_features, self.out_features = lin.in_features, lin.out_features
|
| 22 |
+
self.out_dtype = lin.weight.dtype
|
| 23 |
+
|
| 24 |
+
def forward(self, x):
|
| 25 |
+
shp = x.shape[:-1]
|
| 26 |
+
x2 = x.reshape(-1, self.in_features)
|
| 27 |
+
xq, sx = _quant_rowwise(x2)
|
| 28 |
+
y = torch._scaled_mm(xq, self.wq.t(), scale_a=sx, scale_b=self.sw_t, bias=self.bias, out_dtype=self.out_dtype)
|
| 29 |
+
return y.reshape(*shp, self.out_features)
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def convert_to_fp8(model, skip=("lm_head",), min_dim=1024):
|
| 33 |
+
"""Replace nn.Linear (with in/out >= min_dim) by FP8Linear in place. Returns count."""
|
| 34 |
+
n = 0
|
| 35 |
+
for name, mod in list(model.named_modules()):
|
| 36 |
+
for cname, child in list(mod.named_children()):
|
| 37 |
+
full = f"{name}.{cname}" if name else cname
|
| 38 |
+
if isinstance(child, nn.Linear) and not any(s in full for s in skip) and min(child.in_features, child.out_features) >= min_dim:
|
| 39 |
+
setattr(mod, cname, FP8Linear(child)); n += 1
|
| 40 |
+
return n
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
if __name__ == "__main__":
|
| 44 |
+
import time
|
| 45 |
+
lin = nn.Linear(2048, 6144, bias=False).cuda().to(torch.bfloat16)
|
| 46 |
+
f8 = FP8Linear(lin)
|
| 47 |
+
x = torch.randn(8192, 2048, device="cuda", dtype=torch.bfloat16)
|
| 48 |
+
ref = lin(x); got = f8(x)
|
| 49 |
+
print("rel err", ((ref.float() - got.float()).abs().mean() / ref.float().abs().mean()).item())
|
| 50 |
+
for f, name in [(lin, "bf16 linear"), (f8, "fp8 linear (eager)"), (torch.compile(f8), "fp8 linear (compiled)")]:
|
| 51 |
+
for _ in range(3): f(x)
|
| 52 |
+
torch.cuda.synchronize(); t = time.time()
|
| 53 |
+
for _ in range(20): f(x)
|
| 54 |
+
torch.cuda.synchronize(); print(f"{name:24s} {(time.time()-t)/20*1000:.3f} ms")
|
decider/infer.py
ADDED
|
@@ -0,0 +1,357 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Usable inference API: typed decisions with probabilities, all from one forward pass.
|
| 2 |
+
|
| 3 |
+
from decider.infer import Decider
|
| 4 |
+
d = Decider("runs/r2_full/model")
|
| 5 |
+
out = d.decide("My card was charged twice for the same purchase.",
|
| 6 |
+
[{"question": "Which department should handle this?", "options": ["billing", "technical", "sales"]},
|
| 7 |
+
{"question": "How urgent is this?", "options": ["low", "medium", "high"]}])
|
| 8 |
+
# -> [{'choice': 'billing', 'confidence': 0.97, 'probs': {...}}, {...}]
|
| 9 |
+
"""
|
| 10 |
+
import logging
|
| 11 |
+
|
| 12 |
+
import torch
|
| 13 |
+
from decider.model import DecisionModel, collate
|
| 14 |
+
from decider.prompt import build, MAX_OPTIONS, resolve_layout, chat_template
|
| 15 |
+
from decider import temperature as TT
|
| 16 |
+
from dataclasses import dataclass
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
@dataclass
|
| 20 |
+
class Q:
|
| 21 |
+
text: str; options: list; gold: int = 0
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
@dataclass
|
| 25 |
+
class Example:
|
| 26 |
+
context: str; qs: list; task: str = "infer"; image: bytes = None
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
logger = logging.getLogger(__name__)
|
| 30 |
+
NEUTRAL_NONE = "not listed here"
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def neutralize_options(options):
|
| 34 |
+
"""The training augmentation used the literal 'none of the above', and the model learned that exact string as an
|
| 35 |
+
abstain signal (it abstains even on clear cases when the string is offered). Any option that reads like it is
|
| 36 |
+
rewritten to a neutral phrasing for the model and mapped back in the output."""
|
| 37 |
+
out, back = [], {}
|
| 38 |
+
for o in options:
|
| 39 |
+
key = o.strip().lower()
|
| 40 |
+
if key.startswith("none of the above") or key in ("none of the above", "none", "n/a", "none of these"):
|
| 41 |
+
out.append(NEUTRAL_NONE); back[NEUTRAL_NONE] = o
|
| 42 |
+
else:
|
| 43 |
+
out.append(o)
|
| 44 |
+
return out, back
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class _NoShuffle: # keep option order as given
|
| 48 |
+
def shuffle(self, x): pass
|
| 49 |
+
def sample(self, xs, k): return xs[:k]
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class CompiledSchema:
|
| 53 |
+
def __init__(self, d, rqs, h, index):
|
| 54 |
+
from decider.systemone import row_types
|
| 55 |
+
self.d, self.rqs, self.h, self.index = d, rqs, h, index
|
| 56 |
+
self.types = row_types(rqs, index) # answer type of every schema row (decider.temperature)
|
| 57 |
+
|
| 58 |
+
def batch(self, states, max_state_tokens=32768):
|
| 59 |
+
from decider.systemone import render_state, assemble
|
| 60 |
+
T = TT.for_types(self.d.T_schema, self.d.T_schema_by_type, self.types)
|
| 61 |
+
probs = self.d._se.score(self.h, [render_state(s) for s in states], temperature=T, max_ctx_tokens=max_state_tokens)
|
| 62 |
+
return [{"model": self.d.name, "answers": assemble(self.rqs, self.index, [p.tolist() for p in pr])} for pr in probs]
|
| 63 |
+
|
| 64 |
+
def __call__(self, state, max_state_tokens=32768):
|
| 65 |
+
return self.batch([state], max_state_tokens)[0]
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
class Decider:
|
| 69 |
+
"""One-pass decisions with automatic CUDA, MPS, or CPU device selection.
|
| 70 |
+
|
| 71 |
+
CUDA uses shape-bucketed graphs by default. MPS defaults to float16 and uses
|
| 72 |
+
the optional MPS patch; CPU defaults to bfloat16. Set ``use_graphs=False``
|
| 73 |
+
for eager execution or debugging.
|
| 74 |
+
"""
|
| 75 |
+
def __init__(self, path, device=None, dtype=None, temperature=None, abstain_below=0.0, use_graphs=None, temperature_by_type=None):
|
| 76 |
+
"""The prompt layout comes from decider_config.json: "layout": "chat" (chat-trained checkpoints) wraps every prompt in the
|
| 77 |
+
tokenizer's chat template (decider.prompt.build_chat); no "layout" key is the plain layout of every earlier model.
|
| 78 |
+
An unknown layout raises ValueError before the weights are loaded.
|
| 79 |
+
|
| 80 |
+
Temperatures (decider.temperature): decider_config.json "temperature" and the optional "temperature_by_type"
|
| 81 |
+
{"choice": T, "noul": T, "score": T}. temperature= overrides "temperature" and switches the config's by-type map off;
|
| 82 |
+
temperature_by_type= sets the map explicitly. An invalid temperature or map key raises ValueError before the weights
|
| 83 |
+
are loaded."""
|
| 84 |
+
if device is None:
|
| 85 |
+
device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
|
| 86 |
+
if dtype is None:
|
| 87 |
+
dtype = torch.float16 if str(device).startswith("mps") else torch.bfloat16
|
| 88 |
+
logger.info("Decider device=%s dtype=%s", device, dtype)
|
| 89 |
+
import json, os
|
| 90 |
+
cfg = {}
|
| 91 |
+
try: # model folder may carry decider_config.json (temperature, flags)
|
| 92 |
+
from huggingface_hub import hf_hub_download
|
| 93 |
+
cfg_path = os.path.join(path, "decider_config.json") if os.path.isdir(path) else hf_hub_download(path, "decider_config.json")
|
| 94 |
+
cfg = json.load(open(cfg_path))
|
| 95 |
+
except Exception:
|
| 96 |
+
pass
|
| 97 |
+
self.layout = resolve_layout(cfg)
|
| 98 |
+
(temperature, self.T_by_type), (T_schema, self.T_schema_by_type) = TT.from_config(cfg, temperature, temperature_by_type)
|
| 99 |
+
self.neutralize_none = bool(cfg.get("neutralize_none", True)) # v4 and earlier learned the literal string as an abstain signal
|
| 100 |
+
if use_graphs is None:
|
| 101 |
+
use_graphs = str(device).startswith("cuda")
|
| 102 |
+
if use_graphs:
|
| 103 |
+
from decider.engine import Engine
|
| 104 |
+
self.eng = Engine(path, device=device, dtype=dtype); self.m = self.eng.m
|
| 105 |
+
else:
|
| 106 |
+
if str(device).startswith("mps"):
|
| 107 |
+
from decider.mps_ops import patch_mps
|
| 108 |
+
patch_mps()
|
| 109 |
+
self.eng = None; self.m = DecisionModel(path, dtype=dtype, grad_ckpt=False).to(device).eval()
|
| 110 |
+
self.chat = chat_template(self.m.tok) if self.layout == "chat" else None # None: plain layout, prompts as in 1.1.x
|
| 111 |
+
self.dev = device; self.T = temperature; self.abstain_below = abstain_below
|
| 112 |
+
self.name = "decider-" + str(cfg.get("version", "dev"))
|
| 113 |
+
self.schema_first = bool(cfg.get("schema_first", False)) and self.eng is not None # default layout. Questions-first (the cacheable one) costs accuracy
|
| 114 |
+
self.T_schema = T_schema # (about 1.5 points on fixed label sets, more elsewhere): opt in with schema()
|
| 115 |
+
self.isolated_levels = bool(cfg.get("isolated_levels", False)) # Score levels judged one per row (v8+)
|
| 116 |
+
self._se = None; self._schemas = {}
|
| 117 |
+
|
| 118 |
+
def _decide_items(self, requests, max_ctx_tokens=1536):
|
| 119 |
+
"""The prompt rows of decide_batch: -> (requests with neutralized options, one item per request)."""
|
| 120 |
+
exs = []
|
| 121 |
+
if self.neutralize_none:
|
| 122 |
+
requests = [(context, [dict(q, options=neutralize_options(q["options"])[0], _back=neutralize_options(q["options"])[1]) for q in qs]) for context, qs in requests]
|
| 123 |
+
for context, qs in requests:
|
| 124 |
+
for q in qs:
|
| 125 |
+
assert 2 <= len(q["options"]) <= MAX_OPTIONS, f"2..{MAX_OPTIONS} options required"
|
| 126 |
+
exs.append(Example(context, [Q(q["question"], list(q["options"]), 0) for q in qs], "infer"))
|
| 127 |
+
items = [build(e, self.m.tok, _NoShuffle(), max_options=MAX_OPTIONS, max_ctx_tokens=max_ctx_tokens, chat=self.chat) for e in exs]
|
| 128 |
+
for it, (_, qs) in zip(items, requests): # answer type per slot: a plain question is "choice"
|
| 129 |
+
it["types"] = [q.get("_type", "choice") for q in qs]
|
| 130 |
+
return requests, items
|
| 131 |
+
|
| 132 |
+
@torch.no_grad()
|
| 133 |
+
def decide_batch(self, requests, max_ctx_tokens=1536):
|
| 134 |
+
"""requests: list of (context:str, questions:list[dict(question, options)]). One forward pass for everything."""
|
| 135 |
+
requests, items = self._decide_items(requests, max_ctx_tokens)
|
| 136 |
+
T = TT.for_items(self.T, self.T_by_type, items) # the scalar self.T when there is no by-type map
|
| 137 |
+
if self.eng is not None:
|
| 138 |
+
probs = torch.cat(self.eng.score_items(items, temperature=T))
|
| 139 |
+
else:
|
| 140 |
+
b = collate(items, self.m.tok.pad_token_id)
|
| 141 |
+
logits = self.m.slot_logits(b["input_ids"].to(self.dev), b["attention_mask"].to(self.dev), b["slot_idx"].to(self.dev),
|
| 142 |
+
b["slot_batch"].to(self.dev), b["nopts"].to(self.dev))
|
| 143 |
+
probs = TT.scaled_softmax(logits, TT.slot_temperatures(T, items)).cpu()
|
| 144 |
+
out, k = [], 0
|
| 145 |
+
for context, qs in requests:
|
| 146 |
+
res = []
|
| 147 |
+
for q in qs:
|
| 148 |
+
p = probs[k, :len(q["options"])].tolist(); k += 1
|
| 149 |
+
j = max(range(len(p)), key=p.__getitem__); back = q.get("_back", {})
|
| 150 |
+
names = [back.get(o, o) for o in q["options"]]
|
| 151 |
+
res.append(dict(choice=names[j] if p[j] >= self.abstain_below else None, confidence=p[j],
|
| 152 |
+
probs={o: pi for o, pi in zip(names, p)}, probs_list=p))
|
| 153 |
+
out.append(res)
|
| 154 |
+
return out
|
| 155 |
+
|
| 156 |
+
def decide(self, context, questions, **kw):
|
| 157 |
+
return self.decide_batch([(context, questions)], **kw)[0]
|
| 158 |
+
|
| 159 |
+
# ---- Jev-shaped interface (decider.systemone): state + {id: Choice | Score | Noul with criteria}
|
| 160 |
+
# ---- schema cache (v7+): the questions are run once, requests only run the state (decider.schema_engine)
|
| 161 |
+
def schema(self, questions, independent=True, isolated=None, compile=False):
|
| 162 |
+
"""Compile a fixed set of Jev-shaped questions: schema(state) -> answers; schema.batch([state, ...]) -> [answers]."""
|
| 163 |
+
import json
|
| 164 |
+
from decider.schema_engine import SchemaEngine
|
| 165 |
+
from decider.systemone import render_question
|
| 166 |
+
isolated = self.isolated_levels if isolated is None else isolated
|
| 167 |
+
key = (json.dumps(questions, sort_keys=True, ensure_ascii=False), independent, isolated)
|
| 168 |
+
if key not in self._schemas:
|
| 169 |
+
if self._se is None: self._se = SchemaEngine(self.eng, chat=self.chat)
|
| 170 |
+
if len(self._schemas) >= 64: # drop the oldest schema and its graphs
|
| 171 |
+
old = next(iter(self._schemas)); hid = self._schemas.pop(old)[1].id
|
| 172 |
+
for k in [k for k in self._se.graphs if k[0] == hid]: del self._se.graphs[k]
|
| 173 |
+
from decider.systemone import plan_rows
|
| 174 |
+
rqs = {k: render_question(v) for k, v in questions.items()}
|
| 175 |
+
rows, index = plan_rows(rqs, isolated and independent)
|
| 176 |
+
h = self._se.prepare(rows, independent=independent, compile=compile) # compile=True: ~25 s per (batch, length) shape, 1.6x faster after
|
| 177 |
+
self._schemas[key] = (rqs, h, index)
|
| 178 |
+
return CompiledSchema(self, *self._schemas[key])
|
| 179 |
+
|
| 180 |
+
def system_one(self, state, questions, independent=True, max_state_tokens=32768, max_fwd_tokens=65536, layout=None, isolated=None):
|
| 181 |
+
layout = layout or ("schema_first" if self.schema_first else "state_first")
|
| 182 |
+
isolated = (self.isolated_levels if isolated is None else isolated) and independent
|
| 183 |
+
if layout == "schema_first" and self.eng is not None:
|
| 184 |
+
return self.schema(questions, independent, isolated)(state, max_state_tokens)
|
| 185 |
+
"""independent=True scores every question in its own row (state + that question only), so adding, removing or
|
| 186 |
+
reordering questions cannot change any other answer; the state is run once and its cache forked to every
|
| 187 |
+
question (Engine.score_shared). independent=False packs all questions behind one copy of the state in one row
|
| 188 |
+
(later questions can then see earlier question texts)."""
|
| 189 |
+
from decider.systemone import unique_tokens, assemble
|
| 190 |
+
rqs, index, items = self._system_one_items(state, questions, independent, max_state_tokens, layout, isolated)
|
| 191 |
+
flatp = self._system_one_probs(items, layout, max_fwd_tokens, TT.for_items(self.T, self.T_by_type, items))
|
| 192 |
+
return {"model": self.name, "answers": assemble(rqs, index, flatp),
|
| 193 |
+
"usage": {"input_tokens": unique_tokens(items), "output_tokens": 0}}
|
| 194 |
+
|
| 195 |
+
@torch.no_grad()
|
| 196 |
+
def _system_one_probs(self, items, layout, max_fwd_tokens=65536, temperature=None):
|
| 197 |
+
"""Score system_one's uncached rows. temperature: as Engine.score_items takes it (a number, or one entry per item).
|
| 198 |
+
-> one probability list per question row, in row order."""
|
| 199 |
+
if self.eng is not None and len(items) > 1 and layout == "state_first":
|
| 200 |
+
probs = self.eng.score_shared(items, temperature=temperature)
|
| 201 |
+
else:
|
| 202 |
+
probs = []; per = max(1, max_fwd_tokens // max(len(it["ids"]) for it in items))
|
| 203 |
+
for i in range(0, len(items), per):
|
| 204 |
+
T = TT.item_slice(temperature, i, i + per)
|
| 205 |
+
if self.eng is not None:
|
| 206 |
+
probs += self.eng.score_items(items[i:i + per], temperature=T)
|
| 207 |
+
else:
|
| 208 |
+
bt = collate(items[i:i + per], self.m.tok.pad_token_id)
|
| 209 |
+
lg = self.m.slot_logits(*[bt[k].to(self.dev) for k in ("input_ids", "attention_mask", "slot_idx", "slot_batch", "nopts")])
|
| 210 |
+
pr = TT.scaled_softmax(lg, TT.slot_temperatures(T, items[i:i + per])).cpu(); c = 0
|
| 211 |
+
for it in items[i:i + per]:
|
| 212 |
+
probs.append(pr[c:c + len(it["slots"])]); c += len(it["slots"])
|
| 213 |
+
return [p.tolist() for ps in probs for p in ps]
|
| 214 |
+
|
| 215 |
+
def _system_one_items(self, state, questions, independent=True, max_state_tokens=32768, layout=None, isolated=None):
|
| 216 |
+
"""The prompt rows of system_one's uncached path: -> (rendered questions, answer index, items)."""
|
| 217 |
+
from decider.systemone import render_state, render_question, plan_rows, row_types
|
| 218 |
+
layout = layout or ("schema_first" if self.schema_first else "state_first")
|
| 219 |
+
isolated = (self.isolated_levels if isolated is None else isolated) and independent
|
| 220 |
+
ctx = render_state(state); rqs = {k: render_question(v) for k, v in questions.items()}
|
| 221 |
+
opts = (lambda r: neutralize_options(r["options"])[0]) if self.neutralize_none else (lambda r: list(r["options"]))
|
| 222 |
+
flat, index = plan_rows(rqs, isolated)
|
| 223 |
+
rows = [[r] for r in flat] if independent else [flat]
|
| 224 |
+
items = [build(Example(ctx, [Q(r["question"], opts(r), 0) for r in row]), self.m.tok, _NoShuffle(), max_options=MAX_OPTIONS,
|
| 225 |
+
max_ctx_tokens=max_state_tokens, layout=layout, chat=self.chat) for row in rows]
|
| 226 |
+
types = row_types(rqs, index) # answer type per row (decider.temperature)
|
| 227 |
+
for it, ts in zip(items, [[t] for t in types] if independent else [types]):
|
| 228 |
+
it["types"] = ts
|
| 229 |
+
return rqs, index, items
|
| 230 |
+
|
| 231 |
+
# ---- typed schema interface: {question: {"type": "bool"} | {"type": "choice", "options": [...]}
|
| 232 |
+
# | {"type": "scale", "legend": {"0": "none", "1": "low", ...}}}
|
| 233 |
+
SCHEMA_FORM = ('schema: a map {question: field}, each field one of {"type": "choice", "options": [option strings]}, '
|
| 234 |
+
'{"type": "bool"} or {"type": "scale", "legend": [descriptions] or {level number: description}}')
|
| 235 |
+
|
| 236 |
+
@staticmethod
|
| 237 |
+
def _check_schema(schema):
|
| 238 |
+
"""Raise ValueError naming the expected form when `schema` does not have it (the HTTP server turns this into a 422).
|
| 239 |
+
Every schema that 1.1.2 answered stays accepted, except options or a legend given as a bare string (1.1.2 split it into
|
| 240 |
+
characters); an empty schema is answered with {}."""
|
| 241 |
+
import json, math
|
| 242 |
+
form = Decider.SCHEMA_FORM
|
| 243 |
+
if schema is None:
|
| 244 |
+
raise ValueError("schema is required; " + form)
|
| 245 |
+
if not isinstance(schema, dict):
|
| 246 |
+
raise ValueError(f"schema is a {type(schema).__name__}; " + form)
|
| 247 |
+
for qtext, spec in schema.items():
|
| 248 |
+
where = f"schema[{qtext!r}]"
|
| 249 |
+
if not isinstance(spec, dict):
|
| 250 |
+
raise ValueError(f"{where} is a {type(spec).__name__}, not a field object; " + form)
|
| 251 |
+
t = spec.get("type", "choice")
|
| 252 |
+
if t == "choice":
|
| 253 |
+
opts = spec.get("options")
|
| 254 |
+
if not isinstance(opts, (list, tuple, dict)) or not opts or not all(isinstance(o, str) for o in opts):
|
| 255 |
+
raise ValueError(f'{where}: "options" must be a non-empty list of strings; ' + form)
|
| 256 |
+
elif t == "scale":
|
| 257 |
+
leg = spec.get("legend")
|
| 258 |
+
if not isinstance(leg, (list, tuple, dict)) or not leg:
|
| 259 |
+
raise ValueError(f'{where}: "legend" must be a non-empty list or {{level number: description}} map; ' + form)
|
| 260 |
+
try: # the legend is echoed in the answer, which must serialise
|
| 261 |
+
json.dumps(leg, allow_nan=False)
|
| 262 |
+
except (TypeError, ValueError):
|
| 263 |
+
raise ValueError(f'{where}: the "legend" contains a value that is not finite JSON (NaN or Infinity); ' + form)
|
| 264 |
+
if isinstance(leg, dict):
|
| 265 |
+
try:
|
| 266 |
+
ok = all(math.isfinite(float(k)) for k in leg)
|
| 267 |
+
except (TypeError, ValueError, OverflowError):
|
| 268 |
+
ok = False
|
| 269 |
+
if not ok:
|
| 270 |
+
raise ValueError(f'{where}: the keys of a "legend" map must be finite numbers, e.g. {{"0": "none", "1": "low"}}; ' + form)
|
| 271 |
+
elif t != "bool":
|
| 272 |
+
raise ValueError(f"{where}: unknown field type {t!r}; " + form)
|
| 273 |
+
|
| 274 |
+
@staticmethod
|
| 275 |
+
def _schema_to_questions(schema):
|
| 276 |
+
Decider._check_schema(schema)
|
| 277 |
+
qs = []
|
| 278 |
+
for qtext, spec in schema.items():
|
| 279 |
+
t = spec.get("type", "choice")
|
| 280 |
+
if t == "bool":
|
| 281 |
+
qs.append(dict(question=qtext, options=["no", "yes"], _type="noul"))
|
| 282 |
+
elif t == "choice":
|
| 283 |
+
qs.append(dict(question=qtext, options=list(spec["options"]), _type="choice"))
|
| 284 |
+
elif t == "scale":
|
| 285 |
+
leg = spec["legend"]
|
| 286 |
+
keys = sorted(leg, key=lambda k: float(k)) if isinstance(leg, dict) else list(range(len(leg)))
|
| 287 |
+
labels = [f"{k}: {leg[k]}" if isinstance(leg, dict) else f"{i}: {leg[i]}" for i, k in enumerate(keys)]
|
| 288 |
+
qs.append(dict(question=qtext, options=labels, _keys=keys, _legend=leg, _type="score"))
|
| 289 |
+
else:
|
| 290 |
+
raise ValueError(f"unknown field type {t}")
|
| 291 |
+
return qs
|
| 292 |
+
|
| 293 |
+
def decide_json_batch(self, requests, **kw):
|
| 294 |
+
"""requests: list of (context, schema). Returns one dict per context keyed by question."""
|
| 295 |
+
qss = [self._schema_to_questions(schema) for _, schema in requests]
|
| 296 |
+
raw = self.decide_batch([(ctx, qs) for (ctx, _), qs in zip(requests, qss)], **kw)
|
| 297 |
+
out = []
|
| 298 |
+
for (ctx, schema), qs, res in zip(requests, qss, raw):
|
| 299 |
+
o = {}
|
| 300 |
+
for (qtext, spec), q, r in zip(schema.items(), qs, res):
|
| 301 |
+
t = spec.get("type", "choice")
|
| 302 |
+
if t == "bool":
|
| 303 |
+
o[qtext] = {"noul": round(r["probs"]["yes"], 4), "type": "noul"}
|
| 304 |
+
elif t == "choice":
|
| 305 |
+
o[qtext] = {"choice": r["choice"], "confidence": round(r["confidence"], 4), "type": "choice",
|
| 306 |
+
"probabilities": {k: round(v, 4) for k, v in r["probs"].items()}}
|
| 307 |
+
else:
|
| 308 |
+
p = [r["probs"][lab] for lab in q["options"]]
|
| 309 |
+
keys = q["_keys"]; n = len(p)
|
| 310 |
+
score = sum(float(k) * pi for k, pi in zip(keys, p)) # expected level on the legend scale
|
| 311 |
+
j = max(range(n), key=p.__getitem__)
|
| 312 |
+
o[qtext] = {"score": round(score, 2), "confidence": round(p[j], 4), "type": "scale", "legend": q["_legend"],
|
| 313 |
+
"probabilities": {str(keys[i]): round(pi, 4) for i, pi in enumerate(p)}}
|
| 314 |
+
out.append(o)
|
| 315 |
+
return out
|
| 316 |
+
|
| 317 |
+
def decide_json(self, context, schema, **kw):
|
| 318 |
+
return self.decide_json_batch([(context, schema)], **kw)[0]
|
| 319 |
+
|
| 320 |
+
|
| 321 |
+
if __name__ == "__main__":
|
| 322 |
+
import sys, json, time
|
| 323 |
+
d = Decider(sys.argv[1] if len(sys.argv) > 1 else "runs/r1_200k/model")
|
| 324 |
+
demo = [
|
| 325 |
+
("My card was charged twice for the same purchase and I want the extra charge refunded.",
|
| 326 |
+
[{"question": "Which department should handle this?", "options": ["billing", "technical support", "sales"]},
|
| 327 |
+
{"question": "What is the customer's sentiment?", "options": ["angry", "neutral", "happy"]},
|
| 328 |
+
{"question": "Does this need a refund action?", "options": ["no", "yes"]}]),
|
| 329 |
+
("hey can u turn the lights off in the kitchen",
|
| 330 |
+
[{"question": "What is the intent?", "options": ["smart home control", "set alarm", "play music", "none of the above"]},
|
| 331 |
+
{"question": "Is this request toxic?", "options": ["no", "yes"]}]),
|
| 332 |
+
("The quarterly report shows revenue fell 12% while costs rose sharply.",
|
| 333 |
+
[{"question": "What is the financial sentiment?", "options": ["bearish", "neutral", "bullish"]}]),
|
| 334 |
+
]
|
| 335 |
+
t = time.time(); res = d.decide_batch(demo); dt = time.time() - t
|
| 336 |
+
for (ctx, qs), r in zip(demo, res):
|
| 337 |
+
print("\n>>", ctx)
|
| 338 |
+
for q, a in zip(qs, r):
|
| 339 |
+
print(f" {q['question']:45s} -> {a['choice']!s:22s} p={a['confidence']:.2f} " + " ".join(f"{o}:{p:.2f}" for o, p in a['probs'].items()))
|
| 340 |
+
print(f"\n{sum(len(q) for _, q in demo)} decisions in {dt*1000:.0f} ms (one forward pass)")
|
| 341 |
+
schema = {
|
| 342 |
+
"Revenue currently impacted?": {"type": "bool"},
|
| 343 |
+
"What business impact?": {"type": "choice", "options": ["none", "degraded", "outage"]},
|
| 344 |
+
"Integration issue present?": {"type": "bool"},
|
| 345 |
+
"Account health status?": {"type": "choice", "options": ["healthy", "watch", "at risk"]},
|
| 346 |
+
"Which incident scope?": {"type": "choice", "options": ["single_account", "multi_account", "platform_wide"]},
|
| 347 |
+
"Security concern present?": {"type": "bool"},
|
| 348 |
+
"Duplicate charge reported?": {"type": "bool"},
|
| 349 |
+
"Churn likelihood level?": {"type": "scale", "legend": {"0": "none", "1": "low", "2": "medium", "3": "high"}},
|
| 350 |
+
"Human attention needed?": {"type": "bool"},
|
| 351 |
+
"Immediate feature request?": {"type": "bool"},
|
| 352 |
+
}
|
| 353 |
+
ctx = ("Hi, since this morning our Stripe webhook integration stopped firing and our checkout is down for all customers. "
|
| 354 |
+
"We are losing orders every minute and our partner launch is on Thursday. Also I think we got billed twice last week. "
|
| 355 |
+
"If this is not fixed today we will have to look at other providers.")
|
| 356 |
+
t = time.time(); js = d.decide_json(ctx, schema); dt = time.time() - t
|
| 357 |
+
print(f"\n>> {ctx[:80]}...\n" + json.dumps(js, indent=1)[:3000]); print(f"{len(schema)} typed fields in {dt*1000:.0f} ms (one forward pass)")
|
decider/metrics.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
def ece(conf, correct, bins=15):
|
| 5 |
+
conf = np.asarray(conf); correct = np.asarray(correct, dtype=float)
|
| 6 |
+
edges = np.linspace(0, 1, bins + 1); e = 0.0
|
| 7 |
+
for lo, hi in zip(edges[:-1], edges[1:]):
|
| 8 |
+
m = (conf > lo) & (conf <= hi)
|
| 9 |
+
if m.any():
|
| 10 |
+
e += m.mean() * abs(conf[m].mean() - correct[m].mean())
|
| 11 |
+
return float(e)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def aurc(conf, correct):
|
| 15 |
+
"""Area under risk-coverage curve (lower is better)."""
|
| 16 |
+
order = np.argsort(-np.asarray(conf)); c = np.asarray(correct, dtype=float)[order]
|
| 17 |
+
risk = np.cumsum(1 - c) / np.arange(1, len(c) + 1)
|
| 18 |
+
return float(risk.mean())
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def sel_acc(conf, correct, coverage):
|
| 22 |
+
order = np.argsort(-np.asarray(conf)); c = np.asarray(correct, dtype=float)[order]
|
| 23 |
+
n = max(1, int(round(coverage * len(c))))
|
| 24 |
+
return float(c[:n].mean())
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def summarize(probs, golds, nopts):
|
| 28 |
+
"""probs [N,K] (masked entries 0), golds [N], nopts [N]."""
|
| 29 |
+
probs = np.asarray(probs); golds = np.asarray(golds); nopts = np.asarray(nopts)
|
| 30 |
+
pred = probs.argmax(1); conf = probs.max(1); correct = (pred == golds)
|
| 31 |
+
p_gold = probs[np.arange(len(golds)), golds]
|
| 32 |
+
nll = -np.log(np.clip(p_gold, 1e-12, 1)).mean()
|
| 33 |
+
onehot = np.zeros_like(probs); onehot[np.arange(len(golds)), golds] = 1
|
| 34 |
+
brier = ((probs - onehot) ** 2).sum(1).mean()
|
| 35 |
+
return dict(n=int(len(golds)), acc=float(correct.mean()), nll=float(nll), brier=float(brier), ece=ece(conf, correct),
|
| 36 |
+
aurc=aurc(conf, correct), acc_at_80=sel_acc(conf, correct, 0.8), acc_at_50=sel_acc(conf, correct, 0.5),
|
| 37 |
+
chance=float((1.0 / nopts).mean()), mean_conf=float(conf.mean()))
|
decider/model.py
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Backbone -> slot hidden states -> restricted logits over option letters."""
|
| 2 |
+
import torch, torch.nn as nn, torch.nn.functional as F
|
| 3 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 4 |
+
from decider.prompt import letter_ids, MAX_OPTIONS
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
class DecisionModel(nn.Module):
|
| 8 |
+
def __init__(self, name, dtype=torch.bfloat16, grad_ckpt=True):
|
| 9 |
+
super().__init__()
|
| 10 |
+
if hasattr(torch.backends.cuda, "enable_cudnn_sdp"): # as decider.engine.set_attention_backend_policy: the cuDNN SDPA
|
| 11 |
+
torch.backends.cuda.enable_cudnn_sdp(False) # backend is wrong for masked attention on Blackwell (torch 2.14)
|
| 12 |
+
self.tok = AutoTokenizer.from_pretrained(name)
|
| 13 |
+
self.lm = AutoModelForCausalLM.from_pretrained(name, dtype=dtype)
|
| 14 |
+
if grad_ckpt:
|
| 15 |
+
self.lm.gradient_checkpointing_enable()
|
| 16 |
+
self.register_buffer("letters", torch.tensor(letter_ids(self.tok)), persistent=False)
|
| 17 |
+
|
| 18 |
+
def slot_logits(self, input_ids, attention_mask, slot_idx, slot_batch, nopts):
|
| 19 |
+
"""input_ids [B,T]; slot_idx/slot_batch [N] flat slot positions; nopts [N].
|
| 20 |
+
Returns [N, MAX_OPTIONS] logits with invalid options masked to -inf."""
|
| 21 |
+
h = self.lm.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
|
| 22 |
+
hs = h[slot_batch, slot_idx] # [N,H]
|
| 23 |
+
W = self.lm.lm_head.weight[self.letters] # [K,H]
|
| 24 |
+
logits = F.linear(hs, W).float() # [N,K]
|
| 25 |
+
ar = torch.arange(MAX_OPTIONS, device=logits.device)[None, :]
|
| 26 |
+
logits = logits.masked_fill(ar >= nopts[:, None], float("-inf"))
|
| 27 |
+
return logits
|
| 28 |
+
|
| 29 |
+
def forward(self, batch):
|
| 30 |
+
return self.slot_logits(batch["input_ids"], batch["attention_mask"], batch["slot_idx"], batch["slot_batch"], batch["nopts"])
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def collate(items, pad_id):
|
| 34 |
+
"""items: list of dicts from prompt.build (+ 'task', 'ex_id'). Right-pad."""
|
| 35 |
+
T = max(len(it["ids"]) for it in items)
|
| 36 |
+
T = ((T + 63) // 64) * 64 # few distinct shapes -> fewer kernel (re)compiles
|
| 37 |
+
B = len(items)
|
| 38 |
+
input_ids = torch.full((B, T), pad_id, dtype=torch.long)
|
| 39 |
+
attn = torch.zeros((B, T), dtype=torch.long)
|
| 40 |
+
slot_idx, slot_batch, golds, nopts, tasks, qidx = [], [], [], [], [], []
|
| 41 |
+
for b, it in enumerate(items):
|
| 42 |
+
n = len(it["ids"])
|
| 43 |
+
input_ids[b, :n] = torch.tensor(it["ids"])
|
| 44 |
+
attn[b, :n] = 1
|
| 45 |
+
for k, s in enumerate(it["slots"]):
|
| 46 |
+
slot_idx.append(s); slot_batch.append(b); golds.append(it["golds"][k]); nopts.append(it["nopts"][k])
|
| 47 |
+
tasks.append(it.get("task", "")); qidx.append(k)
|
| 48 |
+
return dict(input_ids=input_ids, attention_mask=attn, slot_idx=torch.tensor(slot_idx), slot_batch=torch.tensor(slot_batch),
|
| 49 |
+
golds=torch.tensor(golds), nopts=torch.tensor(nopts), tasks=tasks, qidx=qidx)
|
decider/mps_moe.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Two MPS slow paths in transformers' Qwen3.5-MoE reference code, replaced in-process (decider-35b-a3b on Apple Silicon).
|
| 2 |
+
|
| 3 |
+
Contributed by @nassersala in issue #6 under the project's Apache-2.0 licence; adapted for the package (out-of-range expert
|
| 4 |
+
ids are dropped as `histc` drops them, and the replacements are installed from `decider.mps_ops.patch_mps`).
|
| 5 |
+
|
| 6 |
+
On an M-series Mac (torch 2.14, transformers 5.17) `torch.histc` takes about 45 ms a call on MPS (once per MoE layer, to
|
| 7 |
+
count tokens per expert) and `torch.linalg.solve_triangular` about 17 ms (twice per linear-attention layer): most of a
|
| 8 |
+
2.8-4 s decision. With both replaced the 35B takes 0.23-0.5 s a decision for typical inputs (reported in issue #6).
|
| 9 |
+
The `histc` replacement is an exact count. The block inverse is within 1e-6 of the MPS solver on real systems; on six items
|
| 10 |
+
the probabilities moved at most 0.033, less than switching to the exact CPU solver does (0.042), and no argmax changed.
|
| 11 |
+
|
| 12 |
+
The replacements are visible only inside those two transformers modules (install() gives each its own view of `torch`), and
|
| 13 |
+
there they act only on MPS tensors of the matching call shape; everything else goes to torch.
|
| 14 |
+
"""
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
_histc, _solve = torch.histc, torch.linalg.solve_triangular
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def count_ids(x, bins, min=0):
|
| 21 |
+
"""histc for integer-valued ids binned one per bin: a count of ids min..min+bins-1; other values are dropped as histc
|
| 22 |
+
drops them (transformers passes expert-parallel sentinels >= num_experts and relies on that)."""
|
| 23 |
+
idx = x.long() - int(min)
|
| 24 |
+
keep = (idx >= 0) & (idx < bins)
|
| 25 |
+
out = torch.zeros(bins, device=x.device, dtype=x.dtype)
|
| 26 |
+
return out.scatter_add_(0, idx.clamp(0, bins - 1), keep.to(x.dtype))
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def histc(input, bins=100, min=0, max=0, *, out=None):
|
| 30 |
+
# expert ids 0..n-1 into n bins over [0, n-1]: bin k holds id k exactly, so this is a count (histc: ~45 ms on MPS)
|
| 31 |
+
if input.device.type == "mps" and input.dim() == 1 and out is None and min == 0 and max > 0 and bins == max - min + 1:
|
| 32 |
+
return count_ids(input, bins, min)
|
| 33 |
+
return _histc(input, bins=bins, min=min, max=max) if out is None else _histc(input, bins=bins, min=min, max=max, out=out)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def unit_lower_inverse(A):
|
| 37 |
+
"""Inverse of unit lower-triangular A (last dim a power of two) by block doubling:
|
| 38 |
+
[[A11, 0], [A21, A22]]^-1 = [[X11, 0], [-X22 A21 X11, X22]]."""
|
| 39 |
+
n = A.shape[-1]
|
| 40 |
+
inv = torch.ones(*A.shape[:-2], n, 1, 1, device=A.device, dtype=A.dtype) # n diagonal 1x1 blocks
|
| 41 |
+
b = 1
|
| 42 |
+
while b < n:
|
| 43 |
+
nb = n // (2 * b)
|
| 44 |
+
blocks = A.reshape(*A.shape[:-2], nb, 2 * b, nb, 2 * b).diagonal(dim1=-4, dim2=-2).movedim(-1, -3)
|
| 45 |
+
x11, x22, a21 = inv[..., 0::2, :, :], inv[..., 1::2, :, :], blocks[..., b:, :b]
|
| 46 |
+
new = torch.zeros(*A.shape[:-2], nb, 2 * b, 2 * b, device=A.device, dtype=A.dtype)
|
| 47 |
+
new[..., :b, :b], new[..., b:, b:], new[..., b:, :b] = x11, x22, -(x22 @ a21 @ x11)
|
| 48 |
+
inv, b = new, 2 * b
|
| 49 |
+
return inv[..., 0, :, :]
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def solve_unit_lower(A, B):
|
| 53 |
+
"""solve_triangular(A, B, upper=False, unitriangular=True): only the strict lower triangle of A is read."""
|
| 54 |
+
n = A.shape[-1]
|
| 55 |
+
return unit_lower_inverse(A.tril(-1) + torch.eye(n, device=A.device, dtype=A.dtype)) @ B
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def solve_triangular(A, B, *, upper, left=True, unitriangular=False, out=None):
|
| 59 |
+
# unit lower-triangular: invert by block doubling, all batched matmuls (the MPS solver: ~17 ms per call)
|
| 60 |
+
n = A.shape[-1]
|
| 61 |
+
if A.device.type == "mps" and not upper and left and unitriangular and out is None and n > 0 and n & (n - 1) == 0:
|
| 62 |
+
return solve_unit_lower(A, B)
|
| 63 |
+
return _solve(A, B, upper=upper, left=left, unitriangular=unitriangular, out=out)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
class _Proxy:
|
| 67 |
+
"""A module's view of `torch` (or `torch.linalg`) with some attributes replaced; everything else is the real one."""
|
| 68 |
+
def __init__(self, real, **over):
|
| 69 |
+
self._real, self._over = real, over
|
| 70 |
+
|
| 71 |
+
def __getattr__(self, name):
|
| 72 |
+
return self._over[name] if name in self._over else getattr(self._real, name)
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
_installed = {}
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def install():
|
| 79 |
+
"""Replace the two operations inside the two transformers modules that call them, and nowhere else: the grouped-MoE
|
| 80 |
+
expert count in `transformers.integrations.moe` and the gated delta rule in `modeling_qwen3_5_moe` see a `torch` whose
|
| 81 |
+
`histc` / `linalg.solve_triangular` are the replacements; every other caller keeps torch's. -> names patched."""
|
| 82 |
+
import importlib
|
| 83 |
+
targets = {"transformers.integrations.moe": dict(histc=histc),
|
| 84 |
+
"transformers.models.qwen3_5_moe.modeling_qwen3_5_moe": dict(linalg=_Proxy(torch.linalg, solve_triangular=solve_triangular))}
|
| 85 |
+
for name, over in targets.items():
|
| 86 |
+
if name in _installed:
|
| 87 |
+
continue
|
| 88 |
+
try:
|
| 89 |
+
mod = importlib.import_module(name)
|
| 90 |
+
except ImportError:
|
| 91 |
+
continue
|
| 92 |
+
if getattr(mod, "torch", None) is not torch:
|
| 93 |
+
continue # not the layout this was written for: leave it alone
|
| 94 |
+
_installed[name] = mod
|
| 95 |
+
mod.torch = _Proxy(torch, **over)
|
| 96 |
+
return sorted(_installed)
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def uninstall():
|
| 100 |
+
for name, mod in list(_installed.items()):
|
| 101 |
+
mod.torch = torch
|
| 102 |
+
del _installed[name]
|
decider/mps_ops.py
ADDED
|
@@ -0,0 +1,268 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""MPS kernels for Qwen3.5 gated-delta attention.
|
| 2 |
+
|
| 3 |
+
The optimization is measured against Transformers' PyTorch reference fallback for
|
| 4 |
+
this kernel; it is not an end-to-end model speedup claim.
|
| 5 |
+
|
| 6 |
+
The gated-delta and L2-normalization code follows Transformers 5.17's
|
| 7 |
+
``modeling_qwen3_5.py`` under the Apache License 2.0. The MPS inversion and
|
| 8 |
+
backend dispatch are Decider additions.
|
| 9 |
+
"""
|
| 10 |
+
import functools
|
| 11 |
+
import inspect
|
| 12 |
+
import logging
|
| 13 |
+
import warnings
|
| 14 |
+
|
| 15 |
+
import torch, torch.nn.functional as F
|
| 16 |
+
from decider.engine import fused_causal_conv1d_fn
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
# Adapted from transformers 5.17.0's modeling_qwen3_5.py (Apache-2.0).
|
| 20 |
+
# The original license applies to the gated-delta and normalization portions;
|
| 21 |
+
# the MPS inversion and dispatch below are Decider additions.
|
| 22 |
+
# Native Metal Shading Language (MSL) JIT kernel via MLX/Metal
|
| 23 |
+
_metal_invert_kernel = None
|
| 24 |
+
_metal_failure_reported = False
|
| 25 |
+
_compat_warning_reported = False
|
| 26 |
+
try:
|
| 27 |
+
import mlx.core as mx
|
| 28 |
+
import mlx.core.fast as fast
|
| 29 |
+
|
| 30 |
+
_METAL_INVERT_SRC = """
|
| 31 |
+
uint col = thread_position_in_threadgroup.x;
|
| 32 |
+
uint mat_idx = threadgroup_position_in_grid.x;
|
| 33 |
+
threadgroup float s_inv[4096];
|
| 34 |
+
|
| 35 |
+
for (uint r = 0; r < 64; ++r) {
|
| 36 |
+
s_inv[r * 64 + col] = (r == col) ? 1.0f : 0.0f;
|
| 37 |
+
}
|
| 38 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 39 |
+
|
| 40 |
+
uint mat_offset = mat_idx * 4096;
|
| 41 |
+
for (uint r = 1; r < 64; ++r) {
|
| 42 |
+
float sum = 0.0f;
|
| 43 |
+
if (col < r) {
|
| 44 |
+
for (uint k = col; k < r; ++k) {
|
| 45 |
+
sum += L[mat_offset + r * 64 + k] * s_inv[k * 64 + col];
|
| 46 |
+
}
|
| 47 |
+
}
|
| 48 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 49 |
+
if (col < r) {
|
| 50 |
+
s_inv[r * 64 + col] = -sum;
|
| 51 |
+
}
|
| 52 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 53 |
+
}
|
| 54 |
+
for (uint r = 0; r < 64; ++r) {
|
| 55 |
+
out_inv[mat_offset + r * 64 + col] = s_inv[r * 64 + col];
|
| 56 |
+
}
|
| 57 |
+
"""
|
| 58 |
+
_metal_invert_kernel = fast.metal_kernel(
|
| 59 |
+
name="invert_unitriangular_64",
|
| 60 |
+
input_names=["L"],
|
| 61 |
+
output_names=["out_inv"],
|
| 62 |
+
source=_METAL_INVERT_SRC,
|
| 63 |
+
)
|
| 64 |
+
except (ImportError, OSError, RuntimeError, AttributeError):
|
| 65 |
+
logging.getLogger(__name__).debug("MLX Metal kernel unavailable; using PyTorch MPS fallback", exc_info=True)
|
| 66 |
+
_metal_invert_kernel = None
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def l2norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-6) -> torch.Tensor:
|
| 70 |
+
return x * torch.rsqrt((x * x).sum(dim=dim, keepdim=True) + eps)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def fast_invert_unitriangular_64(L: torch.Tensor) -> torch.Tensor:
|
| 74 |
+
"""Invert batches of 64x64 unit lower-triangular matrices."""
|
| 75 |
+
global _metal_invert_kernel, _metal_failure_reported
|
| 76 |
+
if _metal_invert_kernel is not None and L.is_mps and not L.requires_grad:
|
| 77 |
+
assert L.dtype == torch.float32, "Metal inversion requires float32 input"
|
| 78 |
+
try:
|
| 79 |
+
orig_shape = L.shape
|
| 80 |
+
L_flat = L.reshape(-1, 64, 64).contiguous()
|
| 81 |
+
num_m = L_flat.shape[0]
|
| 82 |
+
torch.mps.synchronize()
|
| 83 |
+
L_mx = mx.from_dlpack(torch.to_dlpack(L_flat))
|
| 84 |
+
mx.eval(L_mx)
|
| 85 |
+
out_mx = _metal_invert_kernel(
|
| 86 |
+
inputs=[L_mx],
|
| 87 |
+
grid=(num_m * 64, 1, 1),
|
| 88 |
+
threadgroup=(64, 1, 1),
|
| 89 |
+
output_shapes=[(num_m, 64, 64)],
|
| 90 |
+
output_dtypes=[mx.float32],
|
| 91 |
+
)[0]
|
| 92 |
+
mx.eval(out_mx)
|
| 93 |
+
torch.mps.synchronize()
|
| 94 |
+
return torch.from_dlpack(out_mx).to(L.device).reshape(orig_shape)
|
| 95 |
+
except Exception:
|
| 96 |
+
# Optional acceleration must not turn a recoverable backend issue into
|
| 97 |
+
# a failed request; disable the failing kernel for this process.
|
| 98 |
+
_metal_invert_kernel = None
|
| 99 |
+
if not _metal_failure_reported:
|
| 100 |
+
logging.getLogger(__name__).warning("MLX Metal inversion failed; using the PyTorch MPS fallback", exc_info=True)
|
| 101 |
+
_metal_failure_reported = True
|
| 102 |
+
|
| 103 |
+
N = 64
|
| 104 |
+
inv = torch.eye(N, device=L.device, dtype=L.dtype).expand_as(L).clone()
|
| 105 |
+
|
| 106 |
+
# b = 1: 32 blocks of 2x2. For [[1, 0], [l, 1]], inverse is [[1, 0], [-l, 1]]
|
| 107 |
+
idx_row = torch.arange(1, 64, 2, device=L.device)
|
| 108 |
+
idx_col = torch.arange(0, 64, 2, device=L.device)
|
| 109 |
+
inv[..., idx_row, idx_col] = -L[..., idx_row, idx_col]
|
| 110 |
+
|
| 111 |
+
# b = 2: 16 blocks of 4x4
|
| 112 |
+
for j in range(16):
|
| 113 |
+
r1, r2, r3 = j * 4, j * 4 + 2, j * 4 + 4
|
| 114 |
+
inv[..., r2:r3, r1:r2] = -inv[..., r2:r3, r2:r3] @ L[..., r2:r3, r1:r2] @ inv[..., r1:r2, r1:r2]
|
| 115 |
+
|
| 116 |
+
# b = 4: 8 blocks of 8x8
|
| 117 |
+
for j in range(8):
|
| 118 |
+
r1, r2, r3 = j * 8, j * 8 + 4, j * 8 + 8
|
| 119 |
+
inv[..., r2:r3, r1:r2] = -inv[..., r2:r3, r2:r3] @ L[..., r2:r3, r1:r2] @ inv[..., r1:r2, r1:r2]
|
| 120 |
+
|
| 121 |
+
# b = 8: 4 blocks of 16x16
|
| 122 |
+
for j in range(4):
|
| 123 |
+
r1, r2, r3 = j * 16, j * 16 + 8, j * 16 + 16
|
| 124 |
+
inv[..., r2:r3, r1:r2] = -inv[..., r2:r3, r2:r3] @ L[..., r2:r3, r1:r2] @ inv[..., r1:r2, r1:r2]
|
| 125 |
+
|
| 126 |
+
# b = 16: 2 blocks of 32x32
|
| 127 |
+
for j in range(2):
|
| 128 |
+
r1, r2, r3 = j * 32, j * 32 + 16, j * 32 + 32
|
| 129 |
+
inv[..., r2:r3, r1:r2] = -inv[..., r2:r3, r2:r3] @ L[..., r2:r3, r1:r2] @ inv[..., r1:r2, r1:r2]
|
| 130 |
+
|
| 131 |
+
# b = 32: 1 block of 64x64
|
| 132 |
+
inv[..., 32:64, 0:32] = -inv[..., 32:64, 32:64] @ L[..., 32:64, 0:32] @ inv[..., 0:32, 0:32]
|
| 133 |
+
return inv
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def mps_chunk_gated_delta_rule(
|
| 137 |
+
query: torch.Tensor,
|
| 138 |
+
key: torch.Tensor,
|
| 139 |
+
value: torch.Tensor,
|
| 140 |
+
g: torch.Tensor,
|
| 141 |
+
beta: torch.Tensor,
|
| 142 |
+
chunk_size: int = 64,
|
| 143 |
+
initial_state: torch.Tensor | None = None,
|
| 144 |
+
output_final_state: bool = False,
|
| 145 |
+
use_qk_l2norm_in_kernel: bool = False,
|
| 146 |
+
**kwargs,
|
| 147 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 148 |
+
"""Optimized chunk-gated delta rule for Apple Silicon (MPS)."""
|
| 149 |
+
assert chunk_size == 64, f"MPS optimized delta rule requires chunk_size=64, got {chunk_size}"
|
| 150 |
+
initial_dtype = query.dtype
|
| 151 |
+
batch_size, sequence_length, _, k_head_dim = key.shape
|
| 152 |
+
num_v_heads, v_head_dim = value.shape[-2:]
|
| 153 |
+
recurrent_state_shape = (batch_size, num_v_heads, k_head_dim, v_head_dim)
|
| 154 |
+
padded_output_shape = (batch_size, num_v_heads, -1, v_head_dim)
|
| 155 |
+
decay = g
|
| 156 |
+
|
| 157 |
+
query, key, value, beta, decay = [
|
| 158 |
+
x.transpose(1, 2).to(torch.float32, memory_format=torch.contiguous_format)
|
| 159 |
+
for x in (query, key, value, beta, decay)
|
| 160 |
+
]
|
| 161 |
+
if use_qk_l2norm_in_kernel:
|
| 162 |
+
query = l2norm(query, dim=-1, eps=1e-6)
|
| 163 |
+
key = l2norm(key, dim=-1, eps=1e-6)
|
| 164 |
+
scaling = query.shape[-1] ** -0.5
|
| 165 |
+
query = query * scaling
|
| 166 |
+
|
| 167 |
+
pad_size = (chunk_size - sequence_length % chunk_size) % chunk_size
|
| 168 |
+
query = F.pad(query, (0, 0, 0, pad_size))
|
| 169 |
+
key = F.pad(key, (0, 0, 0, pad_size))
|
| 170 |
+
value = F.pad(value, (0, 0, 0, pad_size))
|
| 171 |
+
beta = F.pad(beta, (0, pad_size))
|
| 172 |
+
decay = F.pad(decay, (0, pad_size))
|
| 173 |
+
|
| 174 |
+
v_beta = value * beta.unsqueeze(-1)
|
| 175 |
+
k_beta = key * beta.unsqueeze(-1)
|
| 176 |
+
|
| 177 |
+
query, key, k_beta, v_beta = [
|
| 178 |
+
x.reshape(x.shape[0], x.shape[1], -1, chunk_size, x.shape[-1])
|
| 179 |
+
for x in (query, key, k_beta, v_beta)
|
| 180 |
+
]
|
| 181 |
+
decay = decay.reshape(decay.shape[0], decay.shape[1], -1, chunk_size)
|
| 182 |
+
|
| 183 |
+
strictly_upper_mask = torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=query.device).triu(1)
|
| 184 |
+
cum_decay = decay.cumsum(dim=3)
|
| 185 |
+
pairwise_decay = (cum_decay.unsqueeze(4) - cum_decay.unsqueeze(3)).masked_fill(strictly_upper_mask, float("-inf")).exp()
|
| 186 |
+
|
| 187 |
+
ut_system = (k_beta @ key.transpose(-1, -2)) * pairwise_decay
|
| 188 |
+
intra_chunk_attn = (query @ key.transpose(-1, -2)) * pairwise_decay
|
| 189 |
+
decayed_k_beta = k_beta * cum_decay.exp().unsqueeze(-1)
|
| 190 |
+
|
| 191 |
+
# Fast block divide-and-conquer unitriangular inverse on MPS
|
| 192 |
+
L = ut_system.tril(-1)
|
| 193 |
+
inv = fast_invert_unitriangular_64(L)
|
| 194 |
+
new_values = inv @ v_beta
|
| 195 |
+
k_cumdecay = inv @ decayed_k_beta
|
| 196 |
+
|
| 197 |
+
if initial_state is None:
|
| 198 |
+
last_recurrent_state = torch.zeros(recurrent_state_shape, dtype=new_values.dtype, device=new_values.device)
|
| 199 |
+
else:
|
| 200 |
+
last_recurrent_state = initial_state.to(new_values)
|
| 201 |
+
core_attn_out = torch.zeros_like(new_values)
|
| 202 |
+
|
| 203 |
+
query = query * cum_decay.exp().unsqueeze(-1)
|
| 204 |
+
key = key * (cum_decay[..., -1:] - cum_decay).exp().unsqueeze(-1)
|
| 205 |
+
chunk_decay = cum_decay[..., -1].exp()[..., None, None]
|
| 206 |
+
|
| 207 |
+
num_chunks = query.shape[2]
|
| 208 |
+
qk = torch.cat([query, k_cumdecay], dim=3)
|
| 209 |
+
kt = key.transpose(-1, -2)
|
| 210 |
+
|
| 211 |
+
for i in range(num_chunks):
|
| 212 |
+
qk_state = qk[:, :, i] @ last_recurrent_state
|
| 213 |
+
inter_chunk_attn = qk_state[:, :, :chunk_size]
|
| 214 |
+
v_new = new_values[:, :, i] - qk_state[:, :, chunk_size:]
|
| 215 |
+
core_attn_out[:, :, i] = inter_chunk_attn + intra_chunk_attn[:, :, i] @ v_new
|
| 216 |
+
last_recurrent_state = last_recurrent_state * chunk_decay[:, :, i] + kt[:, :, i] @ v_new
|
| 217 |
+
|
| 218 |
+
last_recurrent_state = None if not output_final_state else last_recurrent_state
|
| 219 |
+
core_attn_out = core_attn_out.reshape(padded_output_shape)[:, :, :sequence_length]
|
| 220 |
+
core_attn_out = core_attn_out.transpose(1, 2).to(initial_dtype, memory_format=torch.contiguous_format)
|
| 221 |
+
return core_attn_out, last_recurrent_state
|
| 222 |
+
|
| 223 |
+
|
| 224 |
+
def patch_mps():
|
| 225 |
+
"""Patch Qwen3.5 operations on MPS tensors while preserving other backends."""
|
| 226 |
+
global _compat_warning_reported
|
| 227 |
+
if not torch.backends.mps.is_available():
|
| 228 |
+
return False
|
| 229 |
+
import os
|
| 230 |
+
if os.environ.get("DECIDER_MPS_MOE_PATCH", "1") != "0": # histc + unit-triangular solve replacements (issue #6, @nassersala)
|
| 231 |
+
from decider import mps_moe
|
| 232 |
+
mps_moe.install()
|
| 233 |
+
try:
|
| 234 |
+
import transformers
|
| 235 |
+
import transformers.models.qwen3_5.modeling_qwen3_5 as mq
|
| 236 |
+
version = tuple(int(part) for part in transformers.__version__.split(".")[:2])
|
| 237 |
+
if not all(callable(getattr(mq, name, None)) for name in ("torch_chunk_gated_delta_rule", "causal_conv1d_fn")):
|
| 238 |
+
raise AttributeError("unsupported Transformers Qwen3.5 implementation")
|
| 239 |
+
params = inspect.signature(mq.torch_chunk_gated_delta_rule).parameters
|
| 240 |
+
required = {"query", "key", "value", "g", "beta", "chunk_size"}
|
| 241 |
+
if version < (5, 17) or not required.issubset(params):
|
| 242 |
+
if not _compat_warning_reported:
|
| 243 |
+
warnings.warn(f"MPS patch requires Transformers >=5.17 with the Qwen3.5 gated-delta signature; got {transformers.__version__}", RuntimeWarning, stacklevel=2)
|
| 244 |
+
_compat_warning_reported = True
|
| 245 |
+
return False
|
| 246 |
+
if getattr(mq.torch_chunk_gated_delta_rule, "_decider_mps_patch", False):
|
| 247 |
+
return True
|
| 248 |
+
|
| 249 |
+
original_delta, original_conv = mq.torch_chunk_gated_delta_rule, mq.causal_conv1d_fn
|
| 250 |
+
delta_params = tuple(inspect.signature(original_delta).parameters)
|
| 251 |
+
query_index, chunk_index = delta_params.index("query"), delta_params.index("chunk_size")
|
| 252 |
+
|
| 253 |
+
@functools.wraps(original_delta)
|
| 254 |
+
def delta(*args, **kwargs):
|
| 255 |
+
query = args[query_index] if len(args) > query_index else kwargs["query"]
|
| 256 |
+
chunk_size = args[chunk_index] if len(args) > chunk_index else kwargs.get("chunk_size", 64)
|
| 257 |
+
return mps_chunk_gated_delta_rule(*args, **kwargs) if query.is_mps and chunk_size == 64 else original_delta(*args, **kwargs)
|
| 258 |
+
|
| 259 |
+
@functools.wraps(original_conv)
|
| 260 |
+
def conv(hidden_states, *args, **kwargs):
|
| 261 |
+
return fused_causal_conv1d_fn(hidden_states, *args, **kwargs) if hidden_states.is_mps else original_conv(hidden_states, *args, **kwargs)
|
| 262 |
+
|
| 263 |
+
delta._decider_mps_patch = True
|
| 264 |
+
mq.torch_chunk_gated_delta_rule, mq.causal_conv1d_fn = delta, conv
|
| 265 |
+
return True
|
| 266 |
+
except (ImportError, AttributeError, TypeError, ValueError):
|
| 267 |
+
logging.getLogger(__name__).debug("MPS patch unavailable; using Transformers fallback", exc_info=True)
|
| 268 |
+
return False
|
decider/prompt.py
ADDED
|
@@ -0,0 +1,287 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Prompt construction. One context, N typed questions, N answer slots.
|
| 2 |
+
|
| 3 |
+
All N decisions are read from a single forward pass: the logits at each
|
| 4 |
+
"Answer k: (" slot are restricted to the option-letter tokens. No answer
|
| 5 |
+
letters are ever inserted, so slot k sees the context and all questions but
|
| 6 |
+
no earlier answers (the decisions are conditionally independent given input).
|
| 7 |
+
"""
|
| 8 |
+
import random
|
| 9 |
+
|
| 10 |
+
LETTERS = "ABCDEFGHIJ"
|
| 11 |
+
NARROW = len(LETTERS) # <= NARROW options: the original "(A) .. (J)" rendering, tokenized as a string (unchanged since v1)
|
| 12 |
+
MAX_OPTIONS = 255 # width of the label head. > NARROW options: "wide" rendering, one label token per option:
|
| 13 |
+
# A..Z then the first 229 two-letter upper-case strings that are single tokens (AA, AB, ...)
|
| 14 |
+
ABSTAIN_PREFIXES = ("none of the above", "none of these", "not listed", "no suitable", "does not apply", "cannot tell")
|
| 15 |
+
ABSTAIN_EXACT = ("other", "unsure", "something else", "neither of these", "other / not covered")
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def is_abstain_option(o):
|
| 19 |
+
o = o.strip().lower()
|
| 20 |
+
return o.startswith(ABSTAIN_PREFIXES) or o in ABSTAIN_EXACT
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
_LABELS = {}
|
| 24 |
+
_OPT_CACHE = {}
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def _enc_opt(tok, text):
|
| 28 |
+
"""Token ids of ") <option text>" (cached: fixed label sets repeat the same strings millions of times)."""
|
| 29 |
+
key = (id(tok), text)
|
| 30 |
+
v = _OPT_CACHE.get(key)
|
| 31 |
+
if v is None:
|
| 32 |
+
v = tok.encode(f") {text}", add_special_tokens=False)
|
| 33 |
+
if len(_OPT_CACHE) < 2_000_000:
|
| 34 |
+
_OPT_CACHE[key] = v
|
| 35 |
+
return v
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def label_table(tok):
|
| 39 |
+
"""(label strings, label token ids), MAX_OPTIONS entries; the first NARROW are A..J so narrow questions are unchanged."""
|
| 40 |
+
key = id(tok)
|
| 41 |
+
if key not in _LABELS:
|
| 42 |
+
import string
|
| 43 |
+
U = string.ascii_uppercase
|
| 44 |
+
names = list(U) + [a + b for a in U for b in U]
|
| 45 |
+
out = []
|
| 46 |
+
for n in names:
|
| 47 |
+
t = tok.encode(n, add_special_tokens=False)
|
| 48 |
+
if len(t) == 1:
|
| 49 |
+
out.append((n, t[0]))
|
| 50 |
+
if len(out) == MAX_OPTIONS:
|
| 51 |
+
break
|
| 52 |
+
assert len(out) == MAX_OPTIONS and len({i for _, i in out}) == MAX_OPTIONS
|
| 53 |
+
_LABELS[key] = ([n for n, _ in out], [i for _, i in out],
|
| 54 |
+
tok.encode("\n(", add_special_tokens=False))
|
| 55 |
+
return _LABELS[key]
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def _select(q, rng, max_options):
|
| 59 |
+
opts = list(range(len(q.options)))
|
| 60 |
+
if len(opts) > max_options:
|
| 61 |
+
# always keep the gold and any abstain-style option (its mere presence must not carry information)
|
| 62 |
+
forced = {q.gold} | {i for i, o in enumerate(q.options) if is_abstain_option(o)}
|
| 63 |
+
others = [i for i in opts if i not in forced]
|
| 64 |
+
opts = rng.sample(others, max_options - len(forced)) + list(forced)
|
| 65 |
+
rng.shuffle(opts)
|
| 66 |
+
return opts
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
def _options_ids(tok, q, opts):
|
| 70 |
+
if len(opts) <= NARROW:
|
| 71 |
+
return tok.encode("".join(f"\n({LETTERS[j]}) {q.options[oi]}" for j, oi in enumerate(opts)), add_special_tokens=False)
|
| 72 |
+
_, lab_ids, open_ids = label_table(tok); out = []
|
| 73 |
+
for j, oi in enumerate(opts):
|
| 74 |
+
out += open_ids + [lab_ids[j]] + _enc_opt(tok, q.options[oi])
|
| 75 |
+
return out
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def build_schema_first(example, tok, rng=None, max_options=NARROW, max_ctx_tokens=1536, chat=None):
|
| 79 |
+
"""Schema-first layout: all question/option blocks, then the context, then one answer slot per question.
|
| 80 |
+
|
| 81 |
+
Question 1: ...\nOptions:\n(A) ... <- prefix: depends only on the questions, so its cache (attention KV and
|
| 82 |
+
\n\nQuestion 2: ... delta-net states) is computed once per schema and reused for every state
|
| 83 |
+
\n\nContext:\n<state>\n\nAnswer 1: (\nAnswer 2: (
|
| 84 |
+
|
| 85 |
+
The three parts are tokenized separately, so `ids[:prefix_len]` is identical for every state.
|
| 86 |
+
chat: a ChatTemplate for a chat-layout model (see `build_chat`); None renders the plain layout above."""
|
| 87 |
+
rng = rng or random
|
| 88 |
+
perms = [_select(q, rng, max_options) for q in example.qs]
|
| 89 |
+
pre = schema_prefix_ids(tok, example.qs, perms, chat=chat)
|
| 90 |
+
suf, slots = schema_suffix_ids(tok, example.context, len(example.qs), max_ctx_tokens, chat=chat)
|
| 91 |
+
return dict(ids=pre + suf, slots=[len(pre) + s for s in slots], golds=[p.index(q.gold) if q.gold in p else -1 for p, q in zip(perms, example.qs)],
|
| 92 |
+
nopts=[len(p) for p in perms], perms=perms, prefix_len=len(pre))
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def schema_prefix_ids(tok, qs, perms=None, chat=None):
|
| 96 |
+
"""Token ids of the question/option blocks (the cacheable part of the schema-first layout).
|
| 97 |
+
chat: a ChatTemplate; the prefix then starts with the template head (the user turn opens before the first question)."""
|
| 98 |
+
multi = len(qs) > 1; pre = [] if chat is None else list(chat.head)
|
| 99 |
+
for k, q in enumerate(qs):
|
| 100 |
+
opts = perms[k] if perms is not None else list(range(len(q.options)))
|
| 101 |
+
pre += tok.encode(f"{chr(10) * 2 if k else ''}Question{' ' + str(k + 1) if multi else ''}: {q.text}\nOptions:", add_special_tokens=False) + _options_ids(tok, q, opts)
|
| 102 |
+
return pre
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def schema_suffix_ids(tok, context, n_q, max_ctx_tokens=1536, chat=None):
|
| 106 |
+
"""Token ids after the schema prefix: the context and one answer slot per question. Returns (ids, slot positions in ids).
|
| 107 |
+
chat: a ChatTemplate; the context then ends the user turn, and the template tail and the answer pieces follow."""
|
| 108 |
+
ids = tok.encode("\n\nContext:\n", add_special_tokens=False) + tok.encode(context, add_special_tokens=False)[:max_ctx_tokens]; slots = []
|
| 109 |
+
if chat is not None:
|
| 110 |
+
ids += chat.tail
|
| 111 |
+
for k in range(n_q):
|
| 112 |
+
ids += chat.answer_ids(tok, k, n_q > 1); slots.append(len(ids) - 1)
|
| 113 |
+
return ids, slots
|
| 114 |
+
for k in range(n_q):
|
| 115 |
+
ids += tok.encode(f"{chr(10) * 2 if k == 0 else chr(10)}Answer{' ' + str(k + 1) if n_q > 1 else ''}: (", add_special_tokens=False); slots.append(len(ids) - 1)
|
| 116 |
+
return ids, slots
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
def build(example, tok, rng=None, max_options=NARROW, max_ctx_tokens=1536, layout="state_first", chat=None):
|
| 120 |
+
"""Returns dict(ids=list[int], slots=list[int], golds=list[int], nopts=list[int], perms=list[list[int]]).
|
| 121 |
+
chat: a ChatTemplate (chat_template(tok)) for a model trained in the chat layout (decider_config.json "layout": "chat");
|
| 122 |
+
None, the default, is the plain layout every earlier model uses, unchanged."""
|
| 123 |
+
if layout == "schema_first":
|
| 124 |
+
return build_schema_first(example, tok, rng, max_options, max_ctx_tokens, chat=chat)
|
| 125 |
+
if chat is not None:
|
| 126 |
+
return build_chat(example, tok, chat, rng, max_options, max_ctx_tokens)
|
| 127 |
+
rng = rng or random
|
| 128 |
+
ctx_ids = tok.encode("Context:\n" + example.context, add_special_tokens=False)[:max_ctx_tokens]
|
| 129 |
+
ids = list(ctx_ids)
|
| 130 |
+
slots, golds, nopts, perms = [], [], [], []
|
| 131 |
+
multi = len(example.qs) > 1
|
| 132 |
+
for k, q in enumerate(example.qs):
|
| 133 |
+
opts = list(range(len(q.options)))
|
| 134 |
+
if len(opts) > max_options:
|
| 135 |
+
# always keep the gold and any abstain-style option (its mere presence must not carry information)
|
| 136 |
+
forced = {q.gold} | {i for i, o in enumerate(q.options) if is_abstain_option(o)}
|
| 137 |
+
others = [i for i in opts if i not in forced]
|
| 138 |
+
keep = rng.sample(others, max_options - len(forced)) + list(forced)
|
| 139 |
+
opts = keep
|
| 140 |
+
rng.shuffle(opts)
|
| 141 |
+
head = f"\n\nQuestion{' ' + str(k + 1) if multi else ''}: {q.text}\nOptions:"
|
| 142 |
+
tail = f"\nAnswer{' ' + str(k + 1) if multi else ''}: ("
|
| 143 |
+
if len(opts) <= NARROW:
|
| 144 |
+
lines = [head] + [f"\n({LETTERS[j]}) {q.options[oi]}" for j, oi in enumerate(opts)] + [tail]
|
| 145 |
+
piece = tok.encode("".join(lines), add_special_tokens=False)
|
| 146 |
+
else: # wide: "\n(" + <label token> + ") text", built from ids so every label is one token
|
| 147 |
+
_, lab_ids, open_ids = label_table(tok)
|
| 148 |
+
piece = tok.encode(head, add_special_tokens=False)
|
| 149 |
+
for j, oi in enumerate(opts):
|
| 150 |
+
piece += open_ids + [lab_ids[j]] + _enc_opt(tok, q.options[oi])
|
| 151 |
+
piece += tok.encode(tail, add_special_tokens=False)
|
| 152 |
+
ids.extend(piece)
|
| 153 |
+
slots.append(len(ids) - 1) # position of " (" token
|
| 154 |
+
golds.append(opts.index(q.gold) if q.gold in opts else -1)
|
| 155 |
+
nopts.append(len(opts))
|
| 156 |
+
perms.append(opts)
|
| 157 |
+
return dict(ids=ids, slots=slots, golds=golds, nopts=nopts, perms=perms)
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
# ---- chat layout (1.2.0) ----------------------------------------------------------------------------------------------
|
| 161 |
+
# A model whose decider_config.json has "layout": "chat" (or "chat_template": true) was trained with every prompt wrapped
|
| 162 |
+
# in its tokenizer's chat template: one user turn holding the plain content, thinking switched off, and the answer pieces in
|
| 163 |
+
# the assistant turn. For Qwen3.5, state-first:
|
| 164 |
+
#
|
| 165 |
+
# <|im_start|>user\n Context:\n<state> \n\nQuestion: <q>\nOptions:\n(A) ..\n(B) ..
|
| 166 |
+
# <|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n Answer: (
|
| 167 |
+
#
|
| 168 |
+
# With several questions in one row all question blocks come first and the answer pieces ("Answer 1: (", "\nAnswer 2: (",
|
| 169 |
+
# ...) follow the template tail. Schema-first: the user turn holds the question blocks and then "\n\nContext:\n<state>".
|
| 170 |
+
# Every part is tokenized on its own (head, context, each question header, each option block, tail, each answer piece);
|
| 171 |
+
# the letter is read at the final " (" token as in the plain layout.
|
| 172 |
+
LAYOUTS = ("plain", "chat")
|
| 173 |
+
|
| 174 |
+
|
| 175 |
+
def resolve_layout(cfg):
|
| 176 |
+
"""The prompt layout a decider_config.json names: "plain" (no "layout" key and no "chat_template": true, which is every
|
| 177 |
+
released model) or "chat". Raises ValueError for any other value, so a model trained in a layout
|
| 178 |
+
this version does not know is refused instead of being read in the wrong one."""
|
| 179 |
+
cfg = cfg or {}
|
| 180 |
+
layout = cfg.get("layout")
|
| 181 |
+
if layout is None:
|
| 182 |
+
layout = "chat" if cfg.get("chat_template") is True else "plain"
|
| 183 |
+
if layout not in LAYOUTS:
|
| 184 |
+
raise ValueError(f"decider_config.json names the prompt layout {layout!r}; this version of decider-ai knows "
|
| 185 |
+
f"{', '.join(repr(x) for x in LAYOUTS)}. Upgrade decider-ai, or check the model's config.")
|
| 186 |
+
if layout == "plain" and cfg.get("chat_template") is True:
|
| 187 |
+
raise ValueError('decider_config.json says "layout": "plain" and "chat_template": true; these contradict each other')
|
| 188 |
+
return layout
|
| 189 |
+
|
| 190 |
+
|
| 191 |
+
class ChatTemplate:
|
| 192 |
+
"""Token ids of the tokenizer's chat template around one user turn, and the answer pieces of the assistant turn.
|
| 193 |
+
|
| 194 |
+
head/tail come from tok.apply_chat_template([user turn], add_generation_prompt=True) with thinking disabled
|
| 195 |
+
(enable_thinking=False where the template knows the switch; a thinking block the template leaves open is closed here),
|
| 196 |
+
with no system prompt. This is the computation of the chat-layout research server the chat models were trained against."""
|
| 197 |
+
SENTINEL = "@@DECIDER_USER_CONTENT@@"
|
| 198 |
+
|
| 199 |
+
def __init__(self, tok, answer="Answer"):
|
| 200 |
+
if not getattr(tok, "chat_template", None):
|
| 201 |
+
raise ValueError('the model is configured for the chat layout ("layout": "chat") but its tokenizer has no chat template')
|
| 202 |
+
msgs = [{"role": "user", "content": self.SENTINEL}]
|
| 203 |
+
kw = {"enable_thinking": False} if "enable_thinking" in tok.chat_template else {}
|
| 204 |
+
s = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, **kw)
|
| 205 |
+
head, tail = s.split(self.SENTINEL)
|
| 206 |
+
for opened, closed in (("<think>\n", "</think>\n\n"), ("<think>", "</think>\n\n"), ("<|channel>thought\n", "<channel|>")):
|
| 207 |
+
if tail.endswith(opened):
|
| 208 |
+
tail += closed
|
| 209 |
+
self.head_text, self.tail_text = head, tail
|
| 210 |
+
self.head = tok.encode(head, add_special_tokens=False)
|
| 211 |
+
self.tail = tok.encode(tail, add_special_tokens=False)
|
| 212 |
+
self.answer = answer
|
| 213 |
+
names, ids, _ = label_table(tok)
|
| 214 |
+
pre = tok.encode(f"{answer}: (", add_special_tokens=False)
|
| 215 |
+
for n, i in list(zip(names, ids))[:NARROW]:
|
| 216 |
+
if tok.encode(f"{answer}: ({n}", add_special_tokens=False) != pre + [i]:
|
| 217 |
+
raise ValueError(f"chat layout: the label {n!r} does not tokenize as one token after {answer!r}: (")
|
| 218 |
+
if tok.encode(tail + f"{answer}: (", add_special_tokens=False) != self.tail + pre:
|
| 219 |
+
raise ValueError("chat layout: the chat template tail merges with the answer prefix")
|
| 220 |
+
|
| 221 |
+
def answer_ids(self, tok, k, multi):
|
| 222 |
+
"""Token ids of answer piece k: "Answer: (" (or "Answer k: (" with several questions), later pieces after a newline."""
|
| 223 |
+
return tok.encode(f"{'' if k == 0 else chr(10)}{self.answer}{' ' + str(k + 1) if multi else ''}: (", add_special_tokens=False)
|
| 224 |
+
|
| 225 |
+
|
| 226 |
+
_CHAT = {}
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
def chat_template(tok):
|
| 230 |
+
"""The ChatTemplate of a tokenizer, cached per tokenizer object."""
|
| 231 |
+
key = id(tok)
|
| 232 |
+
if key not in _CHAT:
|
| 233 |
+
_CHAT[key] = (tok, ChatTemplate(tok)) # keep tok alive so its id is not reused
|
| 234 |
+
return _CHAT[key][1]
|
| 235 |
+
|
| 236 |
+
|
| 237 |
+
def chat_for(tok, cfg):
|
| 238 |
+
"""The ChatTemplate to pass as `chat=` for a model with this decider_config.json, or None for a plain-layout model."""
|
| 239 |
+
return chat_template(tok) if resolve_layout(cfg) == "chat" else None
|
| 240 |
+
|
| 241 |
+
|
| 242 |
+
def load_decider_config(path):
|
| 243 |
+
"""decider_config.json of a model folder or a Hub repository id; {} when the model has none (a raw base model)."""
|
| 244 |
+
import json, os
|
| 245 |
+
if os.path.isdir(path):
|
| 246 |
+
f = os.path.join(path, "decider_config.json")
|
| 247 |
+
return json.load(open(f)) if os.path.exists(f) else {}
|
| 248 |
+
try:
|
| 249 |
+
from huggingface_hub import hf_hub_download
|
| 250 |
+
return json.load(open(hf_hub_download(path, "decider_config.json")))
|
| 251 |
+
except Exception:
|
| 252 |
+
return {}
|
| 253 |
+
|
| 254 |
+
|
| 255 |
+
def chat_for_model(path, tok):
|
| 256 |
+
"""chat_for(tok, the model's decider_config.json): what every script that builds prompts for `path` passes as `chat=`."""
|
| 257 |
+
return chat_for(tok, load_decider_config(path))
|
| 258 |
+
|
| 259 |
+
|
| 260 |
+
def build_chat(example, tok, chat, rng=None, max_options=NARROW, max_ctx_tokens=1536):
|
| 261 |
+
"""State-first chat layout (see the comment above LAYOUTS). Same return fields as `build`. max_ctx_tokens caps the
|
| 262 |
+
tokens of "Context:\\n" + state, as in the plain layout; the template head and tail are not counted."""
|
| 263 |
+
rng = rng or random
|
| 264 |
+
qs = example.qs; multi = len(qs) > 1
|
| 265 |
+
ids = list(chat.head) + tok.encode("Context:\n" + example.context, add_special_tokens=False)[:max_ctx_tokens]
|
| 266 |
+
golds, nopts, perms = [], [], []
|
| 267 |
+
for k, q in enumerate(qs):
|
| 268 |
+
opts = _select(q, rng, max_options)
|
| 269 |
+
ids += tok.encode(f"\n\nQuestion{' ' + str(k + 1) if multi else ''}: {q.text}\nOptions:", add_special_tokens=False) + _options_ids(tok, q, opts)
|
| 270 |
+
golds.append(opts.index(q.gold) if q.gold in opts else -1); nopts.append(len(opts)); perms.append(opts)
|
| 271 |
+
ids += chat.tail
|
| 272 |
+
slots = []
|
| 273 |
+
for k in range(len(qs)):
|
| 274 |
+
ids += chat.answer_ids(tok, k, multi); slots.append(len(ids) - 1)
|
| 275 |
+
return dict(ids=ids, slots=slots, golds=golds, nopts=nopts, perms=perms)
|
| 276 |
+
|
| 277 |
+
|
| 278 |
+
def letter_ids(tok):
|
| 279 |
+
ids = label_table(tok)[1]
|
| 280 |
+
for j, L in enumerate(LETTERS):
|
| 281 |
+
assert tok.encode(L, add_special_tokens=False) == [ids[j]], L
|
| 282 |
+
return ids
|
| 283 |
+
|
| 284 |
+
|
| 285 |
+
def render(example, tok, **kw):
|
| 286 |
+
b = build(example, tok, **kw)
|
| 287 |
+
return tok.decode(b["ids"])
|
decider/prompt_fast.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Row construction for the systemone path that tokenizes the state once per request instead of once per question.
|
| 2 |
+
|
| 3 |
+
`prompt.build` encodes `"Context:\\n" + state` for every row it is called with, so an independent request with n
|
| 4 |
+
questions tokenizes the whole state n times. The question block is encoded as its own string there and simply appended,
|
| 5 |
+
so the ids are exactly `ctx_ids + question_piece`; this module reuses that fact and shares `ctx_ids` across the rows.
|
| 6 |
+
|
| 7 |
+
`build_rows` returns the same item dicts as
|
| 8 |
+
[prompt.build(Example(ctx, [Q(text, options, 0) for text, options in row]), tok, <no-shuffle rng>,
|
| 9 |
+
max_options=MAX_OPTIONS, max_ctx_tokens=max_ctx_tokens) for row in rows]
|
| 10 |
+
(ids, slots, golds, nopts, perms) and is checked against it in tests/test_prompt_fast.py.
|
| 11 |
+
"""
|
| 12 |
+
from decider.prompt import LETTERS, NARROW, MAX_OPTIONS, label_table, _enc_opt, _options_ids
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def context_ids(tok, context, max_ctx_tokens=32768):
|
| 16 |
+
return tok.encode("Context:\n" + context, add_special_tokens=False)[:max_ctx_tokens]
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def question_piece(tok, text, options, k=0, multi=False):
|
| 20 |
+
"""Token ids of one `\\n\\nQuestion...: <text>\\nOptions:...\\nAnswer...: (` block, independent of the state."""
|
| 21 |
+
head = f"\n\nQuestion{' ' + str(k + 1) if multi else ''}: {text}\nOptions:"
|
| 22 |
+
tail = f"\nAnswer{' ' + str(k + 1) if multi else ''}: ("
|
| 23 |
+
if len(options) <= NARROW:
|
| 24 |
+
return tok.encode(head + "".join(f"\n({LETTERS[j]}) {o}" for j, o in enumerate(options)) + tail,
|
| 25 |
+
add_special_tokens=False)
|
| 26 |
+
_, lab_ids, open_ids = label_table(tok)
|
| 27 |
+
piece = tok.encode(head, add_special_tokens=False)
|
| 28 |
+
for j, o in enumerate(options):
|
| 29 |
+
piece += open_ids + [lab_ids[j]] + _enc_opt(tok, o)
|
| 30 |
+
return piece + tok.encode(tail, add_special_tokens=False)
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class _Opts:
|
| 34 |
+
def __init__(self, options): self.options = options
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def _chat_row(tok, chat, ctx, row):
|
| 38 |
+
"""One chat-layout row after the shared ids `ctx` (template head + context): every question block, the template tail,
|
| 39 |
+
then the answer pieces. Equal to prompt.build_chat for the same row (tests/test_layout.py)."""
|
| 40 |
+
multi = len(row) > 1
|
| 41 |
+
ids = list(ctx)
|
| 42 |
+
for k, (text, options) in enumerate(row):
|
| 43 |
+
ids += tok.encode(f"\n\nQuestion{' ' + str(k + 1) if multi else ''}: {text}\nOptions:", add_special_tokens=False)
|
| 44 |
+
ids += _options_ids(tok, _Opts(options), list(range(len(options))))
|
| 45 |
+
ids += chat.tail
|
| 46 |
+
slots = []
|
| 47 |
+
for k in range(len(row)):
|
| 48 |
+
ids += chat.answer_ids(tok, k, multi); slots.append(len(ids) - 1)
|
| 49 |
+
return ids, slots
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def build_rows(tok, context, rows, max_ctx_tokens=32768, chat=None):
|
| 53 |
+
"""rows: list of rows, each a list of (question text, options). -> (items, len(ctx_ids)).
|
| 54 |
+
|
| 55 |
+
Option order is kept as given (no shuffling, no subsetting): systemone.render_question already caps a choice at
|
| 56 |
+
MAX_OPTIONS options, so prompt.build's sampling branch is unreachable here.
|
| 57 |
+
chat: a ChatTemplate (prompt.chat_template) for a chat-layout model; the shared ids are then the template head plus the
|
| 58 |
+
context, and each row is prompt.build_chat's rendering."""
|
| 59 |
+
ctx = context_ids(tok, context, max_ctx_tokens)
|
| 60 |
+
if chat is not None:
|
| 61 |
+
ctx = list(chat.head) + ctx
|
| 62 |
+
items = []
|
| 63 |
+
for row in rows:
|
| 64 |
+
multi = len(row) > 1
|
| 65 |
+
ids = list(ctx); slots = []; nopts = []
|
| 66 |
+
for k, (text, options) in enumerate(row):
|
| 67 |
+
if not 2 <= len(options) <= MAX_OPTIONS:
|
| 68 |
+
raise ValueError(f"2..{MAX_OPTIONS} options required")
|
| 69 |
+
if chat is None:
|
| 70 |
+
ids.extend(question_piece(tok, text, options, k, multi))
|
| 71 |
+
slots.append(len(ids) - 1)
|
| 72 |
+
nopts.append(len(options))
|
| 73 |
+
if chat is not None:
|
| 74 |
+
ids, slots = _chat_row(tok, chat, ctx, row)
|
| 75 |
+
items.append(dict(ids=ids, slots=slots, golds=[0] * len(row), nopts=nopts,
|
| 76 |
+
perms=[list(range(len(o))) for _, o in row]))
|
| 77 |
+
return items, len(ctx)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def unique_tokens(items, ctx_len):
|
| 81 |
+
"""systemone.unique_tokens(items) without re-walking the shared context (all rows start with the same ctx_len ids).
|
| 82 |
+
No rows (a request with no questions) counts 0, as systemone.unique_tokens does."""
|
| 83 |
+
if not items:
|
| 84 |
+
return 0
|
| 85 |
+
sufs = [it["ids"][ctx_len:] for it in items]
|
| 86 |
+
if len(sufs) < 2:
|
| 87 |
+
return ctx_len + sum(len(s) for s in sufs)
|
| 88 |
+
lcp = 0; short = min(len(s) for s in sufs); s0 = sufs[0]
|
| 89 |
+
while lcp < short and all(s[lcp] == s0[lcp] for s in sufs): lcp += 1
|
| 90 |
+
return ctx_len + lcp + sum(len(s) - lcp for s in sufs)
|
decider/schema_engine.py
ADDED
|
@@ -0,0 +1,150 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Schema cache: compute a question schema once, then score states against it.
|
| 2 |
+
|
| 3 |
+
In production the questions are fixed and only the state changes. With the schema-first prompt layout
|
| 4 |
+
(prompt.build_schema_first) the question/option blocks are a prefix that does not depend on the state, so their
|
| 5 |
+
cache - attention K/V for the 6 full-attention layers, conv + recurrent state for the 18 delta-net layers - is computed
|
| 6 |
+
once (`prepare`). A request then runs only "Context: <state>" plus one answer slot per question, as a CUDA graph per
|
| 7 |
+
(batch, length) bucket. The prefix cache is read-only during a request (nothing is written back), so one copy serves
|
| 8 |
+
every batch and every graph.
|
| 9 |
+
|
| 10 |
+
se = SchemaEngine(engine); h = se.prepare([{"question": ..., "options": [...]}, ...])
|
| 11 |
+
probs = se.score(h, ["state 1", "state 2", ...]) # list of [n_questions, MAX_OPTIONS] tensors
|
| 12 |
+
"""
|
| 13 |
+
import time, types, torch, torch.nn.functional as F
|
| 14 |
+
from decider.prompt import schema_prefix_ids, schema_suffix_ids, MAX_OPTIONS
|
| 15 |
+
from decider.engine import read_slots, fill_ids
|
| 16 |
+
|
| 17 |
+
TS_BUCKETS = [32, 48, 64, 96, 128, 192, 256, 384, 512, 768, 1024]
|
| 18 |
+
B_BUCKETS = [1, 2, 4, 8, 16, 32, 64]
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
class _Q:
|
| 22 |
+
def __init__(self, text, options): self.text, self.options = text, options
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class PrefixCache:
|
| 26 |
+
"""Duck-typed transformers Cache over fixed, read-only prefixes, for one suffix forward pass.
|
| 27 |
+
A handle holds P prefixes (P = 1: all questions packed in one prefix; P = n_questions: one prefix per question, so every
|
| 28 |
+
question is scored independently). A batch of R states has R * P rows; row r * P + p continues prefix p."""
|
| 29 |
+
def __init__(self, h, R):
|
| 30 |
+
rep = (lambda t: t.expand(R, *t.shape[1:])) if h.P == 1 else (lambda t: t.repeat(R, *([1] * (t.dim() - 1))))
|
| 31 |
+
self.tp = h.tpmax; self.k = {i: rep(k) for i, k in h.k.items()}; self.v = {i: rep(v) for i, v in h.v.items()}
|
| 32 |
+
self.conv = {i: rep(c).contiguous() for i, c in h.conv.items()}
|
| 33 |
+
self.layers = {i: types.SimpleNamespace(record_past=False, recurrent_states={0: rep(r).contiguous()}) for i, r in h.rec.items()}
|
| 34 |
+
|
| 35 |
+
def has_previous_state(self, layer_idx=None, state_idx=None): return True
|
| 36 |
+
def get_seq_length(self, *a, **k): return self.tp
|
| 37 |
+
def update(self, key, value, layer_idx, *a, **k): return torch.cat([self.k[layer_idx], key], 2), torch.cat([self.v[layer_idx], value], 2)
|
| 38 |
+
def update_conv_state(self, x, layer_idx, **k): return torch.cat([self.conv[layer_idx].to(x.dtype), x], -1)
|
| 39 |
+
def update_recurrent_state(self, s, layer_idx, **k): return s
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class SchemaEngine:
|
| 43 |
+
def __init__(self, engine, use_graphs=True, chat=None):
|
| 44 |
+
"""chat: the ChatTemplate of a chat-layout model (decider.prompt.chat_template); the prefix then starts with the template
|
| 45 |
+
head and the suffix ends with the template tail and the answer pieces (decider.prompt.build_schema_first with chat)."""
|
| 46 |
+
self.chat = chat
|
| 47 |
+
self.e = engine; self.core = engine.core; self.W = engine.W; self.tok = engine.tok; self.dev = engine.dev
|
| 48 |
+
self.use_graphs = use_graphs and engine.use_graphs; self.graphs = {}; self.stats = dict(prepared=0, captures=0, replays=0, eager=0)
|
| 49 |
+
self.compile = bool(engine.cfg.get("compile")); self._compiled = {}
|
| 50 |
+
if self.compile: # every compiled schema graph specialises the model frames again (its cache tensors are constants)
|
| 51 |
+
import torch._dynamo
|
| 52 |
+
torch._dynamo.config.cache_size_limit = 4096; torch._dynamo.config.accumulated_cache_size_limit = 1 << 16
|
| 53 |
+
|
| 54 |
+
@torch.no_grad()
|
| 55 |
+
def prepare(self, questions, independent=False, compile=False):
|
| 56 |
+
"""questions: [{"question": str, "options": [str]}] in the order answers are wanted. Runs the prefix(es) once.
|
| 57 |
+
independent=False: one prefix holding every question (cheapest: a request costs state + n slots).
|
| 58 |
+
independent=True: one prefix per question, one row per question (a request costs n * (state + 1 slot); no question
|
| 59 |
+
can influence another)."""
|
| 60 |
+
qs = [_Q(q["question"], list(q["options"])) for q in questions]
|
| 61 |
+
groups = [[q] for q in qs] if independent else [qs]; pres = [schema_prefix_ids(self.tok, g, chat=self.chat) for g in groups]
|
| 62 |
+
h = types.SimpleNamespace(P=len(groups), nq=len(qs), slots_per_row=1 if independent else len(qs), nopts=[len(q.options) for q in qs], tps=[len(p) for p in pres],
|
| 63 |
+
tpmax=max(len(p) for p in pres), k={}, v={}, conv={}, rec={}, id=self.stats["prepared"],
|
| 64 |
+
compile=bool(compile and self.compile))
|
| 65 |
+
parts = []
|
| 66 |
+
for pre in pres:
|
| 67 |
+
out = self.core(input_ids=torch.tensor(pre, device=self.dev)[None], use_cache=True).past_key_values; d = dict(k={}, v={}, conv={}, rec={})
|
| 68 |
+
for i, layer in enumerate(out.layers):
|
| 69 |
+
if getattr(layer, "recurrent_states", None) is not None and layer.recurrent_states.get(0) is not None:
|
| 70 |
+
d["conv"][i] = layer.conv_states[0]; d["rec"][i] = layer.recurrent_states[0]
|
| 71 |
+
else: # right-pad every prefix's K/V to the longest; the mask hides the padding
|
| 72 |
+
pad = (0, 0, 0, h.tpmax - len(pre)); d["k"][i] = F.pad(layer.keys, pad); d["v"][i] = F.pad(layer.values, pad)
|
| 73 |
+
parts.append(d)
|
| 74 |
+
for name in ("k", "v", "conv", "rec"):
|
| 75 |
+
getattr(h, name).update({i: torch.cat([d[name][i] for d in parts], 0).clone() for i in parts[0][name]})
|
| 76 |
+
self.stats["prepared"] += 1
|
| 77 |
+
return h
|
| 78 |
+
|
| 79 |
+
def _fwd(self, ids, cache, mask, pos):
|
| 80 |
+
hs = self.core(input_ids=ids, past_key_values=cache, attention_mask={"full_attention": mask, "linear_attention": None}, position_ids=pos, use_cache=True).last_hidden_state
|
| 81 |
+
return F.linear(hs, self.W).float()
|
| 82 |
+
|
| 83 |
+
def _static(self, h, R, Ts):
|
| 84 |
+
"""R request slots -> R * P rows. Mask: a row sees its own prefix (not the padding up to tpmax) and the causal suffix."""
|
| 85 |
+
ar = torch.arange(Ts, device=self.dev); tps = torch.tensor(h.tps, device=self.dev).repeat(R) # [R*P]
|
| 86 |
+
pre = (torch.arange(h.tpmax, device=self.dev)[None, :] < tps[:, None])[:, None, None, :].expand(-1, 1, Ts, -1) # [B,1,Ts,tpmax]
|
| 87 |
+
mask = torch.cat([pre, (ar[:, None] >= ar[None, :])[None, None].expand(len(tps), 1, -1, -1)], 3).contiguous()
|
| 88 |
+
return PrefixCache(h, R), mask, (tps[:, None] + ar[None, :]).contiguous()
|
| 89 |
+
|
| 90 |
+
def _capture(self, h, R, Ts):
|
| 91 |
+
B = R * h.P
|
| 92 |
+
ids = torch.full((B, Ts), self.tok.pad_token_id, dtype=torch.long, device=self.dev); cache, mask, pos = self._static(h, R, Ts)
|
| 93 |
+
fwd = self._fwd
|
| 94 |
+
if h.compile: # one compiled function per graph (20-30 s each: only for preloaded schemas): the cache tensors are constants of that graph
|
| 95 |
+
fwd = torch.compile(lambda i: self._fwd(i, cache, mask, pos), dynamic=False)
|
| 96 |
+
call = lambda: fwd(ids)
|
| 97 |
+
else:
|
| 98 |
+
call = lambda: fwd(ids, cache, mask, pos)
|
| 99 |
+
st = torch.cuda.Stream(); st.wait_stream(torch.cuda.current_stream())
|
| 100 |
+
with torch.cuda.stream(st):
|
| 101 |
+
for _ in range(3): call()
|
| 102 |
+
torch.cuda.current_stream().wait_stream(st)
|
| 103 |
+
g = torch.cuda.CUDAGraph()
|
| 104 |
+
with torch.cuda.graph(g, pool=self.e.pool):
|
| 105 |
+
out = call()
|
| 106 |
+
self.stats["captures"] += 1
|
| 107 |
+
return ids, out, g, (cache, mask, pos)
|
| 108 |
+
|
| 109 |
+
def warmup(self, h, batch_sizes=(1, 8, 32), state_tokens=(64, 128, 256)):
|
| 110 |
+
"""Capture (and, for a compiled schema, compile) the graphs for these request-batch sizes and suffix lengths ahead of traffic."""
|
| 111 |
+
t = time.time()
|
| 112 |
+
for R in batch_sizes:
|
| 113 |
+
for Ts in state_tokens:
|
| 114 |
+
Ts = next((x for x in TS_BUCKETS if x >= Ts), TS_BUCKETS[-1])
|
| 115 |
+
if (h.id, R, Ts) not in self.graphs: self.graphs[(h.id, R, Ts)] = self._capture(h, R, Ts)
|
| 116 |
+
torch.cuda.synchronize(); return time.time() - t
|
| 117 |
+
|
| 118 |
+
def tokenize(self, h, context, max_ctx_tokens=1536):
|
| 119 |
+
"""CPU part of a request (do it outside any GPU lock): -> (suffix ids, slot positions)."""
|
| 120 |
+
return schema_suffix_ids(self.tok, context, h.slots_per_row, max_ctx_tokens, chat=self.chat)
|
| 121 |
+
|
| 122 |
+
@staticmethod
|
| 123 |
+
def bucket(n_tokens):
|
| 124 |
+
return next((t for t in TS_BUCKETS if t >= n_tokens), -(-n_tokens // 256) * 256)
|
| 125 |
+
|
| 126 |
+
def score(self, h, contexts, temperature=1.0, max_ctx_tokens=1536):
|
| 127 |
+
"""-> one [n_questions, MAX_OPTIONS] probability tensor per context. temperature: as in score_rows."""
|
| 128 |
+
return self.score_rows(h, [self.tokenize(h, c, max_ctx_tokens) for c in contexts], temperature)
|
| 129 |
+
|
| 130 |
+
@torch.no_grad()
|
| 131 |
+
def score_rows(self, h, rows, temperature=1.0):
|
| 132 |
+
"""rows: [(suffix ids, slots)] from tokenize(). temperature: a number, or a list with one temperature per schema row
|
| 133 |
+
(h.nq values, in the order of prepare's questions), applied to every request."""
|
| 134 |
+
if isinstance(temperature, (list, tuple)) and len(temperature) != h.nq:
|
| 135 |
+
raise ValueError(f"temperature: {len(temperature)} values for a schema with {h.nq} rows")
|
| 136 |
+
Tmax = max(len(r[0]) for r in rows); Ts = next((t for t in TS_BUCKETS if t >= Tmax), None); n = len(rows)
|
| 137 |
+
R = next((b for b in B_BUCKETS if b >= n), n) if Ts else n; Ts = Ts or -(-Tmax // 256) * 256
|
| 138 |
+
ids = fill_ids([x for x, _ in rows for _ in range(h.P)], R * h.P, Ts, self.tok.pad_token_id).to(self.dev, non_blocking=True)
|
| 139 |
+
if self.use_graphs and Ts <= TS_BUCKETS[-1]:
|
| 140 |
+
key = (h.id, R, Ts)
|
| 141 |
+
if key not in self.graphs: self.graphs[key] = self._capture(h, R, Ts)
|
| 142 |
+
s_ids, s_out, g, _ = self.graphs[key]; s_ids.copy_(ids); g.replay(); out = s_out; self.stats["replays"] += 1
|
| 143 |
+
else:
|
| 144 |
+
out = self._fwd(ids, *self._static(h, R, Ts)); self.stats["eager"] += 1
|
| 145 |
+
if h.P == 1: # packed: n slots in one row per request
|
| 146 |
+
rws = [r for r in range(n) for _ in range(h.nq)]; sls = [x for _, sl in rows for x in sl]
|
| 147 |
+
else: # independent: one slot in each of the request's P rows
|
| 148 |
+
rws = [r * h.P + p for r in range(n) for p in range(h.P)]; sls = [sl[0] for _, sl in rows for _ in range(h.P)]
|
| 149 |
+
temps = list(temperature) * n if isinstance(temperature, (list, tuple)) else temperature
|
| 150 |
+
return read_slots(out, rws, sls, h.nopts * n, temps, [h.nq] * n)
|
decider/serve.py
ADDED
|
@@ -0,0 +1,562 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""HTTP server.
|
| 2 |
+
POST /decide {"context": str, "schema": {...}} -> typed JSON decisions (all questions packed in one row)
|
| 3 |
+
schema: {question: {"type": "choice", "options": [...]} | {"type": "bool"} | {"type": "scale", "legend": [...]}};
|
| 4 |
+
a missing or malformed schema is a 422 naming this form (decider.infer.Decider._check_schema)
|
| 5 |
+
POST /v1/systemone {"state": str|object|array, "questions": {id: {...}}} -> the TypeSafe/Jev wire format (decider.systemone):
|
| 6 |
+
Choice (up to 255 described options), Score, Noul; every question is scored in its own row, so answers
|
| 7 |
+
are independent of each other ("independent": false packs them behind one copy of the state instead).
|
| 8 |
+
GET /v1/models /health /stats
|
| 9 |
+
|
| 10 |
+
uvicorn decider.serve:app --host 0.0.0.0 --port 8000 (env: DECIDER_MODEL and the variables below)
|
| 11 |
+
|
| 12 |
+
Execution (1.1.1; docs/SERVING.md has the design and the measurements):
|
| 13 |
+
* decider.engine_v2.EngineV2: CUDA graphs keyed on (batch bucket, padded length bucket) only. The whole grid is captured
|
| 14 |
+
during start-up and the engine is then sealed, so no request pays a graph capture or a torch.compile. Rows longer than the
|
| 15 |
+
last length bucket (8,192 tokens) run eager, in chunks bounded by DECIDER_GRAPH_TOKEN_BUDGET padded tokens.
|
| 16 |
+
* one GPU thread runs every forward; tokenisation runs on a CPU pool; rows waiting at the same moment are partitioned by
|
| 17 |
+
decider.batching.plan_batches, which pads a shorter row into a longer row's bucket when that costs less than a second
|
| 18 |
+
forward (cost model: DECIDER_MERGE_OVERHEAD_TOKENS + rows * padded length). When a collection held rows from more than
|
| 19 |
+
one live request the next one waits up to DECIDER_BATCH_ADAPTIVE_WAIT_MS for more rows, unless its own first row took
|
| 20 |
+
more than 50 ms to arrive (a quiet queue drops the window again); after a single-request collection it does not wait.
|
| 21 |
+
* an independent request with more than one question over a state of at least DECIDER_SHARED_MIN_TOKENS tokens runs the state
|
| 22 |
+
once and forks its cache per question (EngineV2.score_shared, decider.shared_prefix). The fork is made in chunks that fit
|
| 23 |
+
DECIDER_SHARED_FORK_GB, so the peak memory does not grow with the question count. The cuDNN SDPA backend is off
|
| 24 |
+
(decider.engine.set_attention_backend_policy), which is what makes that path agree with the full forward.
|
| 25 |
+
* the schema cache (questions-first layout, prefix cached per question set) is honoured exactly as before: on when
|
| 26 |
+
decider_config.json has "schema_first": true, or with DECIDER_SCHEMA_CACHE=1 on a model trained for that layout. It is
|
| 27 |
+
the one path that does GPU work after start-up that is not a graph replay: the first request with a new schema runs its
|
| 28 |
+
prefix and captures one graph per (schema, batch bucket, state bucket) as in 1.0.x (/stats -> schema_cache.captures).
|
| 29 |
+
Its requests are bounded and admitted like the others, on prefix plus suffix tokens, before any GPU preparation.
|
| 30 |
+
* prompt layout (1.2.0): a model whose decider_config.json has "layout": "chat" (a chat-trained checkpoint) is read in the chat layout
|
| 31 |
+
(decider.prompt.build_chat: the tokenizer's chat template around one user turn, thinking off, "Answer: (" in the assistant
|
| 32 |
+
turn) on every route, including the shared-prefix fork and the schema cache. Every other model is read in the plain
|
| 33 |
+
layout, with the same token ids as 1.1.x. A config naming an unknown layout stops start-up with a ValueError.
|
| 34 |
+
* requests are bounded before they reach the GPU: HTTP 413 when a row, the request or the expanded question count exceeds the
|
| 35 |
+
limits below, HTTP 503 when the outstanding work exceeds DECIDER_MAX_QUEUE_ROWS.
|
| 36 |
+
|
| 37 |
+
Variables (default): DECIDER_DEVICE (auto: cuda, else mps, else cpu) DECIDER_COMPILE (0) DECIDER_FP8 (0) DECIDER_SHARED (1) DECIDER_SHARED_MIN_TOKENS (768)
|
| 38 |
+
DECIDER_SHARED_FORK_GB (8) DECIDER_MERGE_OVERHEAD_TOKENS (512) DECIDER_BATCH_ADAPTIVE_WAIT_MS (2)
|
| 39 |
+
DECIDER_MAX_BATCH (32) DECIDER_BATCH_WAIT_MS (0; DECIDER_MAX_WAIT_MS is an alias) DECIDER_MAX_STATE_TOKENS (32768)
|
| 40 |
+
DECIDER_T_BUCKETS DECIDER_B_BUCKETS DECIDER_GRAPH_TOKEN_BUDGET (32768) DECIDER_WARMUP (1) DECIDER_TOKENIZE_THREADS (8)
|
| 41 |
+
DECIDER_MAX_ROWS (1024) DECIDER_MAX_ROW_TOKENS (DECIDER_MAX_STATE_TOKENS + 4096) DECIDER_MAX_REQUEST_TOKENS (1048576)
|
| 42 |
+
DECIDER_MAX_QUEUE_ROWS (4096) DECIDER_TEMPERATURE DECIDER_SCHEMA_CACHE (0) DECIDER_SCHEMA_MIN_SEEN (2) DECIDER_SCHEMAS
|
| 43 |
+
|
| 44 |
+
Temperatures (1.4.0, decider.temperature): decider_config.json "temperature", and optionally "temperature_by_type"
|
| 45 |
+
{"choice": T, "noul": T, "score": T} (a /decide "bool" field is "noul", a "scale" field is "score"; a missing type uses
|
| 46 |
+
"temperature"). Every row carries the answer type of each of its slots, so a batch that mixes requests and types still
|
| 47 |
+
applies each answer's own temperature. DECIDER_TEMPERATURE replaces "temperature" and switches the by-type map off.
|
| 48 |
+
/health and the ready line report the temperature every answer type gets.
|
| 49 |
+
"""
|
| 50 |
+
import asyncio, json, os, time
|
| 51 |
+
from concurrent.futures import ThreadPoolExecutor
|
| 52 |
+
from contextlib import asynccontextmanager
|
| 53 |
+
from fastapi import FastAPI, HTTPException
|
| 54 |
+
from pydantic import BaseModel
|
| 55 |
+
from decider import systemone as S1
|
| 56 |
+
from decider import temperature as TT
|
| 57 |
+
from decider.batching import DEFAULT_MERGE_OVERHEAD_TOKENS, plan_batches
|
| 58 |
+
from decider.prompt import build, MAX_OPTIONS, resolve_layout, chat_template
|
| 59 |
+
from decider.prompt_fast import build_rows, unique_tokens
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def _env_int(name, default):
|
| 63 |
+
return int(os.environ.get(name, default))
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
MODEL = os.environ.get("DECIDER_MODEL", "runs/r3_v2/model")
|
| 67 |
+
MAX_BATCH = _env_int("DECIDER_MAX_BATCH", 32)
|
| 68 |
+
BATCH_WAIT_MS = float(os.environ.get("DECIDER_BATCH_WAIT_MS", os.environ.get("DECIDER_MAX_WAIT_MS", "0")))
|
| 69 |
+
ADAPTIVE_WAIT_MS = float(os.environ.get("DECIDER_BATCH_ADAPTIVE_WAIT_MS", "2")) # window after a collection that held >1 request; 0 disables
|
| 70 |
+
ADAPTIVE_IDLE_RESET_MS = 50.0 # an idler queue than this drops the adaptive window again
|
| 71 |
+
MERGE_OVERHEAD_TOKENS = _env_int("DECIDER_MERGE_OVERHEAD_TOKENS", DEFAULT_MERGE_OVERHEAD_TOKENS)
|
| 72 |
+
MAX_STATE_TOKENS = _env_int("DECIDER_MAX_STATE_TOKENS", 32768)
|
| 73 |
+
SHARED = os.environ.get("DECIDER_SHARED", "1") == "1"
|
| 74 |
+
SHARED_MIN_TOKENS = _env_int("DECIDER_SHARED_MIN_TOKENS", 768)
|
| 75 |
+
DEVICE = os.environ.get("DECIDER_DEVICE", "auto") # auto | cuda[:i] | mps | cpu
|
| 76 |
+
COMPILE = os.environ.get("DECIDER_COMPILE", "0") == "1"
|
| 77 |
+
FP8 = os.environ.get("DECIDER_FP8", "0") == "1"
|
| 78 |
+
WARMUP = os.environ.get("DECIDER_WARMUP", "1") == "1"
|
| 79 |
+
TOKENIZE_THREADS = _env_int("DECIDER_TOKENIZE_THREADS", 8)
|
| 80 |
+
GRAPH_TOKEN_BUDGET = _env_int("DECIDER_GRAPH_TOKEN_BUDGET", 32768)
|
| 81 |
+
# request bounds, checked after tokenisation and before anything is queued
|
| 82 |
+
MAX_ROWS = _env_int("DECIDER_MAX_ROWS", 1024) # scoring rows per request (questions, expanded score levels)
|
| 83 |
+
MAX_ROW_TOKENS = _env_int("DECIDER_MAX_ROW_TOKENS", MAX_STATE_TOKENS + 4096) # tokens in one row: truncated state + question block
|
| 84 |
+
MAX_REQUEST_TOKENS = _env_int("DECIDER_MAX_REQUEST_TOKENS", 1 << 20) # sum of row lengths of one request
|
| 85 |
+
MAX_QUEUE_ROWS = _env_int("DECIDER_MAX_QUEUE_ROWS", 4096) # rows admitted and not yet scored, over all requests
|
| 86 |
+
DECIDE_MAX_CTX_TOKENS = 1536 # /decide context cap, unchanged from 1.0.x
|
| 87 |
+
|
| 88 |
+
MODEL_NAME = "decider"; TEMP = 1.0; TEMP_SCHEMA = 1.0; TEMP_BY_TYPE = {}; TEMP_SCHEMA_BY_TYPE = {}; RELEASE_DATE = "2026-09-17"; ISOLATED = False; NEUTRALIZE_NONE = True
|
| 89 |
+
SCHEMA_FIRST = False; LAYOUT = "plain"; CHAT = None; se = None; squeue = None; schemas = {}; seen = {}
|
| 90 |
+
eng = None; queue = None; gpu = None; cpu = None; batcher_task = None; schema_task = None
|
| 91 |
+
outstanding = 0; REQ_SEQ = 0
|
| 92 |
+
stats = dict(requests=0, batches=0, decisions=0, rows=0, shared_prefix_requests=0, errors=0, rejected_too_large=0,
|
| 93 |
+
rejected_overloaded=0, batch_hist={}, bucket_hist={})
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def _ints(name):
|
| 97 |
+
v = os.environ.get(name)
|
| 98 |
+
return [int(x) for x in v.split(",")] if v else None
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
def load_config(path):
|
| 102 |
+
"""decider_config.json from a model folder or a Hub repository (the way decider.infer.Decider resolves it)."""
|
| 103 |
+
try:
|
| 104 |
+
if os.path.isdir(path):
|
| 105 |
+
return json.load(open(os.path.join(path, "decider_config.json")))
|
| 106 |
+
from huggingface_hub import hf_hub_download
|
| 107 |
+
return json.load(open(hf_hub_download(path, "decider_config.json")))
|
| 108 |
+
except Exception:
|
| 109 |
+
return {}
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
# ---- request preparation (CPU) -------------------------------------------
|
| 113 |
+
class _NoShuffle:
|
| 114 |
+
def shuffle(self, x): pass
|
| 115 |
+
def sample(self, xs, k): return xs[:k]
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def prepare(tok, state, questions, independent, isolated=False, max_state_tokens=32768, chat=None):
|
| 119 |
+
"""Render, plan the rows, tokenize the state once. -> (rqs, index, items, ctx_len). The rows are those
|
| 120 |
+
`prompt.build` produces for the same request (tests/test_prompt_fast.py, tests/test_serve_prepare.py); chat: the
|
| 121 |
+
ChatTemplate of a chat-layout model, None for the plain layout."""
|
| 122 |
+
ctx = S1.render_state(state)
|
| 123 |
+
rqs = {k: S1.render_question(v) for k, v in questions.items()}
|
| 124 |
+
flat, index = S1.plan_rows(rqs, isolated and independent)
|
| 125 |
+
pairs = [(r["question"], list(r["options"])) for r in flat]
|
| 126 |
+
rows = [[p] for p in pairs] if independent else [pairs]
|
| 127 |
+
items, ctx_len = build_rows(tok, ctx, rows, max_ctx_tokens=max_state_tokens, chat=chat)
|
| 128 |
+
types = S1.row_types(rqs, index) # answer type per slot (decider.temperature)
|
| 129 |
+
for it, ts in zip(items, [[t] for t in types] if independent else [types]):
|
| 130 |
+
it["types"] = ts
|
| 131 |
+
return rqs, index, items, ctx_len
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def _prepare_s1(state, questions, independent):
|
| 135 |
+
return prepare(eng.tok, state, questions, independent, ISOLATED, MAX_STATE_TOKENS, chat=CHAT)
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
def _prepare_decide(context, schema):
|
| 139 |
+
from decider.infer import Decider, Example, Q, neutralize_options
|
| 140 |
+
qs = Decider._schema_to_questions(schema)
|
| 141 |
+
for q in qs:
|
| 142 |
+
if NEUTRALIZE_NONE:
|
| 143 |
+
q["options"], q["_back"] = neutralize_options(q["options"])
|
| 144 |
+
ex = Example(context, [Q(q["question"], list(q["options"]), 0) for q in qs])
|
| 145 |
+
it = build(ex, eng.tok, _NoShuffle(), max_options=MAX_OPTIONS, max_ctx_tokens=DECIDE_MAX_CTX_TOKENS, chat=CHAT)
|
| 146 |
+
it["types"] = [q["_type"] for q in qs] # bool -> noul, scale -> score (decider.temperature)
|
| 147 |
+
return qs, it
|
| 148 |
+
|
| 149 |
+
|
| 150 |
+
def _format_decide(schema, qs, probs):
|
| 151 |
+
o = {}
|
| 152 |
+
for (qtext, spec), q, p in zip(schema.items(), qs, probs):
|
| 153 |
+
p = p[:len(q["options"])].tolist(); t = spec.get("type", "choice"); j = max(range(len(p)), key=p.__getitem__)
|
| 154 |
+
back = q.get("_back", {}); names = [back.get(x, x) for x in q["options"]]
|
| 155 |
+
if t == "bool":
|
| 156 |
+
o[qtext] = {"noul": round(p[1], 4), "type": "noul"}
|
| 157 |
+
elif t == "choice":
|
| 158 |
+
o[qtext] = {"choice": names[j], "confidence": round(p[j], 4), "type": "choice",
|
| 159 |
+
"probabilities": {k: round(v, 4) for k, v in zip(names, p)}}
|
| 160 |
+
else:
|
| 161 |
+
keys = q["_keys"]; score = sum(float(k) * pi for k, pi in zip(keys, p))
|
| 162 |
+
o[qtext] = {"score": round(score, 2), "confidence": round(p[j], 4), "type": "scale", "legend": q["_legend"],
|
| 163 |
+
"probabilities": {str(keys[i]): round(pi, 4) for i, pi in enumerate(p)}}
|
| 164 |
+
return o
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
# ---- limits ---------------------------------------------------------------
|
| 168 |
+
def check_size(row_lengths, max_rows=None, max_row_tokens=None, max_request_tokens=None):
|
| 169 |
+
"""Raise HTTPException(413) when a request exceeds the row count, per-row token or total token limit."""
|
| 170 |
+
max_rows = MAX_ROWS if max_rows is None else max_rows
|
| 171 |
+
max_row_tokens = MAX_ROW_TOKENS if max_row_tokens is None else max_row_tokens
|
| 172 |
+
max_request_tokens = MAX_REQUEST_TOKENS if max_request_tokens is None else max_request_tokens
|
| 173 |
+
n, total, longest = len(row_lengths), sum(row_lengths), max(row_lengths, default=0)
|
| 174 |
+
if n > max_rows:
|
| 175 |
+
msg = f"too many questions: the request expands to {n} scoring rows, the limit is {max_rows} (DECIDER_MAX_ROWS)"
|
| 176 |
+
elif longest > max_row_tokens:
|
| 177 |
+
msg = f"too many tokens: one row has {longest} tokens, the limit is {max_row_tokens} per row (DECIDER_MAX_ROW_TOKENS)"
|
| 178 |
+
elif total > max_request_tokens:
|
| 179 |
+
msg = f"too many tokens: the request has {total} tokens over {n} rows, the limit is {max_request_tokens} (DECIDER_MAX_REQUEST_TOKENS)"
|
| 180 |
+
else:
|
| 181 |
+
return
|
| 182 |
+
stats["rejected_too_large"] += 1
|
| 183 |
+
raise HTTPException(413, msg)
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def _admit(n):
|
| 187 |
+
"""Reserve n rows of outstanding work or raise HTTPException(503)."""
|
| 188 |
+
global outstanding
|
| 189 |
+
if outstanding + n > MAX_QUEUE_ROWS:
|
| 190 |
+
stats["rejected_overloaded"] += 1
|
| 191 |
+
raise HTTPException(503, f"server busy: {outstanding} rows queued, the limit is {MAX_QUEUE_ROWS} (DECIDER_MAX_QUEUE_ROWS); retry later")
|
| 192 |
+
outstanding += n
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
def _release(n):
|
| 196 |
+
global outstanding
|
| 197 |
+
outstanding -= n
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
# ---- GPU work ----------------------------------------------------------------
|
| 201 |
+
def _score_items(items):
|
| 202 |
+
return eng.score_items(items, temperature=TT.for_items(TEMP, TEMP_BY_TYPE, items)) # the scalar TEMP without a by-type map
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
def _score_shared(items):
|
| 206 |
+
return eng.score_shared(items, temperature=TT.for_items(TEMP, TEMP_BY_TYPE, items))
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
async def _collect(q, wait_ms=None, adaptive_ms=0.0, idle_reset_ms=None):
|
| 210 |
+
"""Take what is already queued and go.
|
| 211 |
+
|
| 212 |
+
`wait_ms` (DECIDER_BATCH_WAIT_MS) is the unconditional collection window. `adaptive_ms` is the extra window the
|
| 213 |
+
batcher asks for after a collection that held more than one live request; it is dropped again when the first row of
|
| 214 |
+
this collection took longer than `idle_reset_ms` to arrive, so an isolated request after a quiet period never waits.
|
| 215 |
+
"""
|
| 216 |
+
wait = BATCH_WAIT_MS if wait_ms is None else wait_ms
|
| 217 |
+
reset = ADAPTIVE_IDLE_RESET_MS if idle_reset_ms is None else idle_reset_ms
|
| 218 |
+
t0 = time.monotonic()
|
| 219 |
+
batch = [await q.get()]
|
| 220 |
+
if adaptive_ms > 0 and (time.monotonic() - t0) * 1000 <= reset:
|
| 221 |
+
wait = max(wait, adaptive_ms)
|
| 222 |
+
deadline = time.monotonic() + wait / 1000
|
| 223 |
+
while len(batch) < MAX_BATCH:
|
| 224 |
+
try:
|
| 225 |
+
batch.append(q.get_nowait())
|
| 226 |
+
except asyncio.QueueEmpty:
|
| 227 |
+
timeout = deadline - time.monotonic()
|
| 228 |
+
if timeout <= 0: break
|
| 229 |
+
try: batch.append(await asyncio.wait_for(q.get(), timeout))
|
| 230 |
+
except asyncio.TimeoutError: break
|
| 231 |
+
return batch
|
| 232 |
+
|
| 233 |
+
|
| 234 |
+
def _bucketed(n):
|
| 235 |
+
"""False for a row longer than the engine's last captured length bucket. Those run eager at a request-specific shape,
|
| 236 |
+
so they are grouped by exact padded length as in 1.1.0 instead of being padded into another row's bucket."""
|
| 237 |
+
t_bucket = getattr(eng, "t_bucket", None)
|
| 238 |
+
return True if t_bucket is None else t_bucket(n) is not None
|
| 239 |
+
|
| 240 |
+
|
| 241 |
+
def adaptive_ms(batch):
|
| 242 |
+
"""The extra collection window the next collection may use: DECIDER_BATCH_ADAPTIVE_WAIT_MS when this collection held
|
| 243 |
+
live rows from more than one request, 0 otherwise. Rows whose future is already done (a cancelled or failed request)
|
| 244 |
+
are not counted: they are not evidence that requests are overlapping."""
|
| 245 |
+
if ADAPTIVE_WAIT_MS <= 0:
|
| 246 |
+
return 0.0
|
| 247 |
+
return ADAPTIVE_WAIT_MS if len({rid for fut, _, rid in batch if not fut.done()}) > 1 else 0.0
|
| 248 |
+
|
| 249 |
+
|
| 250 |
+
async def batcher():
|
| 251 |
+
"""One forward per planned group. Rows queued at the same moment are partitioned by decider.batching.plan_batches:
|
| 252 |
+
a shorter row is padded into a longer row's bucket when that costs less than running a second forward."""
|
| 253 |
+
loop = asyncio.get_running_loop()
|
| 254 |
+
extra = 0.0 # the adaptive window the previous collection earned
|
| 255 |
+
while True:
|
| 256 |
+
batch = await _collect(queue, BATCH_WAIT_MS, extra)
|
| 257 |
+
extra = adaptive_ms(batch)
|
| 258 |
+
groups = plan_batches([len(it["ids"]) for _, it, _ in batch], eng.pad_len, eng.max_rows, MAX_BATCH,
|
| 259 |
+
MERGE_OVERHEAD_TOKENS, _bucketed)
|
| 260 |
+
for T, idx in groups:
|
| 261 |
+
part = [batch[i] for i in idx]
|
| 262 |
+
stats["bucket_hist"][T] = stats["bucket_hist"].get(T, 0) + len(part)
|
| 263 |
+
try:
|
| 264 |
+
probs = await loop.run_in_executor(gpu, _score_items, [it for _, it, _ in part])
|
| 265 |
+
for (fut, _, _), p in zip(part, probs):
|
| 266 |
+
if not fut.done(): fut.set_result(p)
|
| 267 |
+
except Exception as e:
|
| 268 |
+
for fut, _, _ in part:
|
| 269 |
+
if not fut.done(): fut.set_exception(e)
|
| 270 |
+
stats["batches"] += 1
|
| 271 |
+
stats["batch_hist"][len(part)] = stats["batch_hist"].get(len(part), 0) + 1
|
| 272 |
+
|
| 273 |
+
|
| 274 |
+
# ---- schema cache (questions-first layout; opt-in) --------------------------------
|
| 275 |
+
def _schema_key(questions, independent):
|
| 276 |
+
return (json.dumps(questions, sort_keys=True, ensure_ascii=False), independent)
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
class _SQ:
|
| 280 |
+
def __init__(self, text, options): self.text, self.options = text, options
|
| 281 |
+
|
| 282 |
+
|
| 283 |
+
def _plan_schema(questions, independent, state):
|
| 284 |
+
"""CPU part of a schema-cache request: prefix token lengths (from the cached handle when the schema is known, otherwise
|
| 285 |
+
tokenised here, without touching the GPU) and the suffix row (context plus answer slots). -> (tps, (suffix ids, slots)).
|
| 286 |
+
Raises ValueError for an invalid question."""
|
| 287 |
+
from decider.prompt import schema_prefix_ids, schema_suffix_ids
|
| 288 |
+
rqs = {k: S1.render_question(v) for k, v in questions.items()}; rows, _ = S1.plan_rows(rqs, ISOLATED and independent)
|
| 289 |
+
cached = schemas.get(_schema_key(questions, independent))
|
| 290 |
+
if cached is not None:
|
| 291 |
+
tps = list(cached[1].tps)
|
| 292 |
+
else:
|
| 293 |
+
qs = [_SQ(x["question"], list(x["options"])) for x in rows]
|
| 294 |
+
tps = [len(schema_prefix_ids(eng.tok, g, chat=CHAT)) for g in ([[q] for q in qs] if independent else [qs])]
|
| 295 |
+
row = schema_suffix_ids(eng.tok, S1.render_state(state), 1 if independent else len(rows), MAX_STATE_TOKENS, chat=CHAT)
|
| 296 |
+
return tps, row
|
| 297 |
+
|
| 298 |
+
|
| 299 |
+
def _schema_handle(questions, independent, compile=False):
|
| 300 |
+
"""Compile (or look up) the question schema: its prefix is run once, requests then only run the state. GPU thread."""
|
| 301 |
+
key = _schema_key(questions, independent)
|
| 302 |
+
if key not in schemas:
|
| 303 |
+
rqs = {k: S1.render_question(v) for k, v in questions.items()}; rows, index = S1.plan_rows(rqs, ISOLATED and independent)
|
| 304 |
+
if len(schemas) >= 128:
|
| 305 |
+
old = next(iter(schemas)); hid = schemas.pop(old)[1].id
|
| 306 |
+
for k in [k for k in se.graphs if k[0] == hid]: del se.graphs[k]
|
| 307 |
+
h = se.prepare(rows, independent=independent, compile=compile)
|
| 308 |
+
h.types = S1.row_types(rqs, index) # answer type per schema row (decider.temperature)
|
| 309 |
+
schemas[key] = (rqs, h, index)
|
| 310 |
+
return schemas[key]
|
| 311 |
+
|
| 312 |
+
|
| 313 |
+
def _worth_caching(questions, independent):
|
| 314 |
+
"""A schema gets a cached prefix and CUDA graphs from its second request on."""
|
| 315 |
+
key = _schema_key(questions, independent)
|
| 316 |
+
if key in schemas: return True
|
| 317 |
+
if len(seen) > 50000: seen.clear()
|
| 318 |
+
seen[key] = seen.get(key, 0) + 1
|
| 319 |
+
return seen[key] >= _env_int("DECIDER_SCHEMA_MIN_SEEN", 2)
|
| 320 |
+
|
| 321 |
+
|
| 322 |
+
def _score_schema(h, rows):
|
| 323 |
+
return se.score_rows(h, rows, temperature=TT.for_types(TEMP_SCHEMA, TEMP_SCHEMA_BY_TYPE, h.types))
|
| 324 |
+
|
| 325 |
+
|
| 326 |
+
async def schema_batcher():
|
| 327 |
+
loop = asyncio.get_running_loop()
|
| 328 |
+
while True:
|
| 329 |
+
batch = await _collect(squeue)
|
| 330 |
+
groups = {}
|
| 331 |
+
for fut, h, row in batch: groups.setdefault((h.id, se.bucket(len(row[0]))), (h, []))[1].append((fut, row))
|
| 332 |
+
for h, items in groups.values():
|
| 333 |
+
step = max(1, MAX_BATCH // h.P)
|
| 334 |
+
for i in range(0, len(items), step):
|
| 335 |
+
chunk = items[i:i + step]
|
| 336 |
+
try:
|
| 337 |
+
probs = await loop.run_in_executor(gpu, _score_schema, h, [c for _, c in chunk])
|
| 338 |
+
for (fut, _), p in zip(chunk, probs):
|
| 339 |
+
if not fut.done(): fut.set_result(p)
|
| 340 |
+
except Exception as e:
|
| 341 |
+
for fut, _ in chunk:
|
| 342 |
+
if not fut.done(): fut.set_exception(e)
|
| 343 |
+
stats["schema_batches"] = stats.get("schema_batches", 0) + 1
|
| 344 |
+
|
| 345 |
+
|
| 346 |
+
# ---- start-up / shutdown -------------------------------------------------------
|
| 347 |
+
def apply_config(cfg):
|
| 348 |
+
global MODEL_NAME, TEMP, TEMP_SCHEMA, TEMP_BY_TYPE, TEMP_SCHEMA_BY_TYPE, RELEASE_DATE, ISOLATED, NEUTRALIZE_NONE, SCHEMA_FIRST, LAYOUT
|
| 349 |
+
LAYOUT = resolve_layout(cfg) # ValueError for an unknown layout, before the engine is built
|
| 350 |
+
NEUTRALIZE_NONE = bool(cfg.get("neutralize_none", True))
|
| 351 |
+
MODEL_NAME = "decider-" + str(cfg.get("version", "dev"))
|
| 352 |
+
# ValueError for a temperature that is not finite and > 0 or an unknown by-type key, before the engine is built.
|
| 353 |
+
# DECIDER_TEMPERATURE replaces "temperature" and switches "temperature_by_type" off (decider.temperature).
|
| 354 |
+
(TEMP, TEMP_BY_TYPE), (TEMP_SCHEMA, TEMP_SCHEMA_BY_TYPE) = TT.from_config(cfg, os.environ.get("DECIDER_TEMPERATURE"))
|
| 355 |
+
RELEASE_DATE = str(cfg.get("release_date", RELEASE_DATE))
|
| 356 |
+
ISOLATED = bool(cfg.get("isolated_levels", False))
|
| 357 |
+
trained = bool(cfg.get("schema_first", False) or cfg.get("schema_first_trained", False))
|
| 358 |
+
SCHEMA_FIRST = trained and (bool(cfg.get("schema_first", False)) or os.environ.get("DECIDER_SCHEMA_CACHE", "0") == "1")
|
| 359 |
+
|
| 360 |
+
|
| 361 |
+
def start_workers(loop=None):
|
| 362 |
+
"""Queues, executors and batcher tasks. Needs `eng` set; a test can set a stand-in engine and call this directly."""
|
| 363 |
+
global queue, gpu, cpu, batcher_task, squeue, schema_task
|
| 364 |
+
loop = loop or asyncio.get_running_loop()
|
| 365 |
+
if gpu is None: gpu = ThreadPoolExecutor(max_workers=1, thread_name_prefix="gpu")
|
| 366 |
+
if cpu is None: cpu = ThreadPoolExecutor(max_workers=TOKENIZE_THREADS, thread_name_prefix="tok")
|
| 367 |
+
queue = asyncio.Queue(); batcher_task = loop.create_task(batcher())
|
| 368 |
+
if SCHEMA_FIRST:
|
| 369 |
+
squeue = asyncio.Queue(); schema_task = loop.create_task(schema_batcher())
|
| 370 |
+
|
| 371 |
+
|
| 372 |
+
def resolve_device(requested=None):
|
| 373 |
+
"""The device and dtype the server runs on. `auto` picks as decider.infer.Decider does: CUDA, else MPS, else CPU; float16
|
| 374 |
+
on MPS, bfloat16 elsewhere. An explicit device that is not available, or FP8 / torch.compile off CUDA, is a start-up error
|
| 375 |
+
that says so, instead of the torch assertion a CUDA call raises on a build without CUDA."""
|
| 376 |
+
import torch
|
| 377 |
+
req = (DEVICE if requested is None else requested).strip().lower()
|
| 378 |
+
if req in ("", "auto"):
|
| 379 |
+
dev = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
|
| 380 |
+
else:
|
| 381 |
+
dev = req
|
| 382 |
+
kind = dev.split(":")[0]
|
| 383 |
+
if kind == "cuda" and not torch.cuda.is_available():
|
| 384 |
+
raise RuntimeError(f"DECIDER_DEVICE={req}, but torch.cuda.is_available() is False on this machine. "
|
| 385 |
+
"Set DECIDER_DEVICE=mps or cpu (or leave it at auto).")
|
| 386 |
+
if kind == "mps" and not torch.backends.mps.is_available():
|
| 387 |
+
raise RuntimeError(f"DECIDER_DEVICE={req}, but torch.backends.mps.is_available() is False on this machine.")
|
| 388 |
+
if kind not in ("cuda", "mps", "cpu"):
|
| 389 |
+
raise RuntimeError(f"DECIDER_DEVICE={req}: expected auto, cuda, cuda:<index>, mps or cpu.")
|
| 390 |
+
if not dev.startswith("cuda") and (FP8 or COMPILE):
|
| 391 |
+
raise RuntimeError(f"DECIDER_FP8 and DECIDER_COMPILE need CUDA; the server is starting on {dev}. Unset them.")
|
| 392 |
+
return dev, (torch.float16 if dev.startswith("mps") else torch.bfloat16)
|
| 393 |
+
|
| 394 |
+
|
| 395 |
+
async def _start():
|
| 396 |
+
global eng, se, gpu, CHAT
|
| 397 |
+
from decider.engine_v2 import EngineV2
|
| 398 |
+
apply_config(load_config(MODEL))
|
| 399 |
+
gpu = ThreadPoolExecutor(max_workers=1, thread_name_prefix="gpu")
|
| 400 |
+
loop = asyncio.get_running_loop()
|
| 401 |
+
dev, dtype = resolve_device()
|
| 402 |
+
eng = await loop.run_in_executor(gpu, lambda: EngineV2(
|
| 403 |
+
MODEL, device=dev, dtype=dtype, compile=COMPILE, fp8=FP8, max_ctx_tokens=MAX_STATE_TOKENS, t_buckets=_ints("DECIDER_T_BUCKETS"),
|
| 404 |
+
b_buckets=_ints("DECIDER_B_BUCKETS"), token_budget=GRAPH_TOKEN_BUDGET))
|
| 405 |
+
CHAT = chat_template(eng.tok) if LAYOUT == "chat" else None
|
| 406 |
+
print("[serve] engine", dict({k: v for k, v in eng.cfg.items() if k not in ("t_buckets", "b_buckets")}, device=dev, layout=LAYOUT), flush=True)
|
| 407 |
+
if SCHEMA_FIRST:
|
| 408 |
+
from decider.schema_engine import SchemaEngine
|
| 409 |
+
se = SchemaEngine(eng, chat=CHAT); print("[serve] schema cache on", flush=True)
|
| 410 |
+
pre = os.environ.get("DECIDER_SCHEMAS") # JSON: [{"questions": {...}, "independent": true, "batch_sizes": [1, 8, 32], "state_tokens": [64, 256]}]
|
| 411 |
+
for spec in (json.load(open(pre)) if pre else []):
|
| 412 |
+
_, h, _ = await loop.run_in_executor(gpu, _schema_handle, spec["questions"], spec.get("independent", True), COMPILE)
|
| 413 |
+
t = await loop.run_in_executor(gpu, se.warmup, h, spec.get("batch_sizes", (1, 8, 32)), spec.get("state_tokens", (64, 128, 256)))
|
| 414 |
+
print(f"[serve] preloaded schema with {h.nq} rows, prefix {sum(h.tps)} tokens, graphs ready in {t:.0f}s", flush=True)
|
| 415 |
+
if WARMUP and eng.use_graphs: # off CUDA there are no graphs to capture; every request runs eager
|
| 416 |
+
t = await loop.run_in_executor(gpu, lambda: eng.warmup(log=lambda s: print(s, flush=True)))
|
| 417 |
+
print(f"[serve] captured {len(eng.graphs)} graphs in {t:.0f}s", flush=True)
|
| 418 |
+
eng.seal()
|
| 419 |
+
print("[serve] ready", json.dumps(dict(model=MODEL_NAME, layout=LAYOUT, temperature=TEMP, temperature_by_type=TT.effective(TEMP, TEMP_BY_TYPE),
|
| 420 |
+
**({"temperature_schema_first_by_type": TT.effective(TEMP_SCHEMA, TEMP_SCHEMA_BY_TYPE)} if SCHEMA_FIRST else {}),
|
| 421 |
+
isolated_levels=ISOLATED, schema_first=SCHEMA_FIRST,
|
| 422 |
+
shared=SHARED, graphs=len(eng.graphs), limits=dict(
|
| 423 |
+
max_rows=MAX_ROWS, max_row_tokens=MAX_ROW_TOKENS, max_request_tokens=MAX_REQUEST_TOKENS,
|
| 424 |
+
max_queue_rows=MAX_QUEUE_ROWS))), flush=True)
|
| 425 |
+
start_workers(loop)
|
| 426 |
+
|
| 427 |
+
|
| 428 |
+
def _stop():
|
| 429 |
+
global gpu, cpu, batcher_task, schema_task
|
| 430 |
+
for t in (batcher_task, schema_task):
|
| 431 |
+
if t is not None: t.cancel()
|
| 432 |
+
for ex in (gpu, cpu):
|
| 433 |
+
if ex is not None: ex.shutdown(wait=False, cancel_futures=True)
|
| 434 |
+
gpu = cpu = batcher_task = schema_task = None
|
| 435 |
+
|
| 436 |
+
|
| 437 |
+
@asynccontextmanager
|
| 438 |
+
async def lifespan(app):
|
| 439 |
+
await _start()
|
| 440 |
+
yield
|
| 441 |
+
_stop()
|
| 442 |
+
|
| 443 |
+
|
| 444 |
+
app = FastAPI(title="decider", lifespan=lifespan)
|
| 445 |
+
|
| 446 |
+
|
| 447 |
+
# ---- routes -------------------------------------------------------------------
|
| 448 |
+
class Req(BaseModel):
|
| 449 |
+
context: str
|
| 450 |
+
schema_: object = None # any JSON value: Decider._check_schema gives the 422 (pydantic's echoes NaN/Infinity and cannot be serialised)
|
| 451 |
+
model_config = {"populate_by_name": True}
|
| 452 |
+
def __init__(self, **kw):
|
| 453 |
+
if "schema" in kw: kw["schema_"] = kw.pop("schema")
|
| 454 |
+
super().__init__(**kw)
|
| 455 |
+
|
| 456 |
+
|
| 457 |
+
class S1Req(BaseModel):
|
| 458 |
+
state: object
|
| 459 |
+
questions: dict
|
| 460 |
+
model: str | None = None
|
| 461 |
+
independent: bool = True
|
| 462 |
+
layout: str | None = None # "state_first" forces the uncached layout on a schema-first model
|
| 463 |
+
|
| 464 |
+
|
| 465 |
+
def _alive():
|
| 466 |
+
return batcher_task is not None and not batcher_task.done() and (schema_task is None or not schema_task.done())
|
| 467 |
+
|
| 468 |
+
|
| 469 |
+
async def _queued(items):
|
| 470 |
+
"""Queue one request's rows. Every row carries the request's id, which is what the batcher's adaptive wait reads."""
|
| 471 |
+
global REQ_SEQ
|
| 472 |
+
loop = asyncio.get_running_loop(); futs = []
|
| 473 |
+
REQ_SEQ += 1; rid = REQ_SEQ
|
| 474 |
+
for it in items:
|
| 475 |
+
f = loop.create_future(); futs.append(f); queue.put_nowait((f, it, rid))
|
| 476 |
+
return await asyncio.gather(*futs)
|
| 477 |
+
|
| 478 |
+
|
| 479 |
+
@app.post("/decide")
|
| 480 |
+
async def decide(r: Req):
|
| 481 |
+
loop = asyncio.get_running_loop()
|
| 482 |
+
try:
|
| 483 |
+
qs, it = await loop.run_in_executor(cpu, _prepare_decide, r.context, r.schema_)
|
| 484 |
+
except (ValueError, KeyError) as e:
|
| 485 |
+
stats["errors"] += 1
|
| 486 |
+
raise HTTPException(422, str(e))
|
| 487 |
+
if not qs: # empty schema: nothing to score (the 1.1.2 answer, without a forward)
|
| 488 |
+
stats["requests"] += 1; return {}
|
| 489 |
+
check_size([len(it["ids"])])
|
| 490 |
+
_admit(1)
|
| 491 |
+
try:
|
| 492 |
+
probs = (await _queued([it]))[0]
|
| 493 |
+
except Exception:
|
| 494 |
+
stats["errors"] += 1; raise
|
| 495 |
+
finally:
|
| 496 |
+
_release(1)
|
| 497 |
+
stats["requests"] += 1; stats["decisions"] += len(qs); stats["rows"] += 1
|
| 498 |
+
return _format_decide(r.schema_, qs, probs)
|
| 499 |
+
|
| 500 |
+
|
| 501 |
+
@app.post("/v1/systemone")
|
| 502 |
+
async def systemone(r: S1Req):
|
| 503 |
+
loop = asyncio.get_running_loop()
|
| 504 |
+
if SCHEMA_FIRST and r.questions and r.layout != "state_first" and _worth_caching(r.questions, r.independent):
|
| 505 |
+
try: # CPU: rows, prefix lengths, suffix ids -> the complete cost before any GPU work
|
| 506 |
+
tps, row = await loop.run_in_executor(cpu, _plan_schema, r.questions, r.independent, r.state)
|
| 507 |
+
except ValueError as e:
|
| 508 |
+
stats["errors"] += 1
|
| 509 |
+
raise HTTPException(422, str(e))
|
| 510 |
+
check_size([tp + len(row[0]) for tp in tps])
|
| 511 |
+
_admit(len(tps))
|
| 512 |
+
try:
|
| 513 |
+
rqs, h, index = await loop.run_in_executor(gpu, _schema_handle, r.questions, r.independent)
|
| 514 |
+
fut = loop.create_future(); await squeue.put((fut, h, row)); p = await fut
|
| 515 |
+
except Exception:
|
| 516 |
+
stats["errors"] += 1; raise
|
| 517 |
+
finally:
|
| 518 |
+
_release(len(tps))
|
| 519 |
+
stats["requests"] += 1; stats["decisions"] += len(rqs); stats["schema_requests"] = stats.get("schema_requests", 0) + 1
|
| 520 |
+
return {"model": MODEL_NAME, "answers": S1.assemble(rqs, index, [pk.tolist() for pk in p]),
|
| 521 |
+
"usage": {"input_tokens": len(row[0]) * h.P, "cached_tokens": sum(h.tps), "output_tokens": 0}}
|
| 522 |
+
try:
|
| 523 |
+
rqs, index, items, ctx_len = await loop.run_in_executor(cpu, _prepare_s1, r.state, r.questions, r.independent)
|
| 524 |
+
except ValueError as e:
|
| 525 |
+
stats["errors"] += 1
|
| 526 |
+
raise HTTPException(422, str(e))
|
| 527 |
+
check_size([len(it["ids"]) for it in items])
|
| 528 |
+
_admit(len(items))
|
| 529 |
+
try:
|
| 530 |
+
if SHARED and len(items) > 1 and min(len(it["ids"]) for it in items) >= SHARED_MIN_TOKENS:
|
| 531 |
+
res = await loop.run_in_executor(gpu, _score_shared, items) # long state: run it once, fork the cache per question
|
| 532 |
+
stats["shared_prefix_requests"] += 1
|
| 533 |
+
else:
|
| 534 |
+
res = await _queued(items)
|
| 535 |
+
except Exception:
|
| 536 |
+
stats["errors"] += 1; raise
|
| 537 |
+
finally:
|
| 538 |
+
_release(len(items))
|
| 539 |
+
probs = [p for ps in res for p in ps] # one prob row per question, request order
|
| 540 |
+
stats["requests"] += 1; stats["decisions"] += len(rqs); stats["rows"] += len(items)
|
| 541 |
+
return {"model": MODEL_NAME, "answers": S1.assemble(rqs, index, [p.tolist() for p in probs]),
|
| 542 |
+
"usage": {"input_tokens": unique_tokens(items, ctx_len), "output_tokens": 0}}
|
| 543 |
+
|
| 544 |
+
|
| 545 |
+
@app.get("/v1/models")
|
| 546 |
+
async def models():
|
| 547 |
+
return {"models": [{"name": MODEL_NAME, "description": "decider: one-pass typed decisions with calibrated probabilities", "release_date": RELEASE_DATE}]}
|
| 548 |
+
|
| 549 |
+
|
| 550 |
+
@app.get("/health")
|
| 551 |
+
async def health():
|
| 552 |
+
return {"ok": eng is not None and bool(getattr(eng, "sealed", True)) and _alive(), "model": MODEL,
|
| 553 |
+
"device": str(getattr(eng, "dev", "")) if eng is not None else None, "layout": LAYOUT,
|
| 554 |
+
"temperature": TEMP, "temperature_by_type": TT.effective(TEMP, TEMP_BY_TYPE),
|
| 555 |
+
**({"temperature_schema_first_by_type": TT.effective(TEMP_SCHEMA, TEMP_SCHEMA_BY_TYPE)} if SCHEMA_FIRST else {})}
|
| 556 |
+
|
| 557 |
+
|
| 558 |
+
@app.get("/stats")
|
| 559 |
+
async def get_stats():
|
| 560 |
+
return dict(stats, outstanding_rows=outstanding, engine=eng.stats if eng else None, graphs=len(eng.graphs) if eng else 0,
|
| 561 |
+
sealed=bool(eng and getattr(eng, "sealed", False)), schema_cache=dict(se.stats, schemas=len(schemas)) if se else None,
|
| 562 |
+
limits=dict(max_rows=MAX_ROWS, max_row_tokens=MAX_ROW_TOKENS, max_request_tokens=MAX_REQUEST_TOKENS, max_queue_rows=MAX_QUEUE_ROWS))
|
decider/shared_prefix.py
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared-prefix scoring with a bounded memory fork, used by both engines.
|
| 2 |
+
|
| 3 |
+
Rows that start with the same tokens (one state, one question per row) run the shared prefix once and then only the
|
| 4 |
+
question suffixes, against a copy of the prefix cache. Until 1.1.0 the copy was made with
|
| 5 |
+
`cache.reorder_cache(zeros(n))`, which expands the prefix to all n rows at once: a 31k-token state with 32 questions
|
| 6 |
+
costs n times the prefix cache, which is 133 GB on a 31B model and is wasteful on the 2B. Here the prefix cache is
|
| 7 |
+
forked in chunks of m rows, m chosen so that one fork fits a byte budget, and the rows are scored chunk by chunk.
|
| 8 |
+
|
| 9 |
+
The budget is `DECIDER_SHARED_FORK_GB` (8 GB), capped at half of the memory currently free on the device (device-free
|
| 10 |
+
plus the caching allocator's reserved-but-unused blocks); `m = clamp(budget // prefix_bytes, 1, n)`. Chunking changes
|
| 11 |
+
which rows share a forward and how far each chunk's suffixes are padded, so answers can move by the usual bf16
|
| 12 |
+
reduction-order amount; they do not depend on it mathematically (right padding, causal layers).
|
| 13 |
+
|
| 14 |
+
Cache layout. transformers 5.17 keeps one layer object per layer in `cache.layers`. An attention layer carries
|
| 15 |
+
`keys` and `values` tensors `[batch, heads, seq, head_dim]`; a sparse-attention layer adds `indexer_keys`
|
| 16 |
+
`[batch, seq, dim]`; a linear-attention layer carries `conv_states` and `recurrent_states` dicts of tensors whose first
|
| 17 |
+
dimension is the batch. Qwen3.5 gives a `DynamicCache` whose 24 layers are a mix of the two objects. `cache_state_tensors`
|
| 18 |
+
enumerates whatever of `STATE_NAMES` is present rather than assuming a layout, and `fork_cache` builds a new cache
|
| 19 |
+
object from them without touching the original (the linear-attention layers update their states with `copy_()`, so a
|
| 20 |
+
fork that shared storage with the prefix would corrupt it for the next chunk).
|
| 21 |
+
|
| 22 |
+
A layout we do not know is not chunked. If a layer holds any other tensor attribute, or a dict of tensors, outside
|
| 23 |
+
`STATE_NAMES`, we cannot tell whether it carries a batch dimension, so `score_shared` takes one fork of all n rows --
|
| 24 |
+
the 1.1.0 behaviour, correct but unbounded -- and counts it in `engine.stats["shared_unchunked_layout"]`.
|
| 25 |
+
|
| 26 |
+
The prefix and suffix forwards are eager, at request-specific shapes, and are correct only with the cuDNN SDPA backend
|
| 27 |
+
off (`decider.engine.set_attention_backend_policy`, applied in both engines' `__init__`).
|
| 28 |
+
"""
|
| 29 |
+
import copy
|
| 30 |
+
import os
|
| 31 |
+
|
| 32 |
+
import torch
|
| 33 |
+
import torch.nn.functional as F
|
| 34 |
+
|
| 35 |
+
from decider.engine import fill_ids, read_slots
|
| 36 |
+
from decider.temperature import item_slice, slot_temperatures
|
| 37 |
+
|
| 38 |
+
DEFAULT_FORK_GB = 8.0
|
| 39 |
+
STATE_NAMES = ("keys", "values", "indexer_keys", "conv_states", "recurrent_states")
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
def common_prefix_len(ids):
|
| 43 |
+
"""Length of the longest common prefix of the rows, capped one token below the shortest row."""
|
| 44 |
+
lcp = 0
|
| 45 |
+
short = min(len(x) for x in ids) - 1
|
| 46 |
+
while lcp < short and all(x[lcp] == ids[0][lcp] for x in ids):
|
| 47 |
+
lcp += 1
|
| 48 |
+
return lcp
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def _get(container, key):
|
| 52 |
+
return container[key] if isinstance(container, dict) else getattr(container, key)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _set(container, key, value):
|
| 56 |
+
if isinstance(container, dict):
|
| 57 |
+
container[key] = value
|
| 58 |
+
else:
|
| 59 |
+
setattr(container, key, value)
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
def cache_state_tensors(cache):
|
| 63 |
+
"""Every tensor in the cache whose first dimension is the batch, as (container, key) pairs.
|
| 64 |
+
|
| 65 |
+
Covers `STATE_NAMES`: the attention layers' `keys`/`values`, a sparse-attention layer's `indexer_keys` and the
|
| 66 |
+
linear-attention layers' `conv_states`/`recurrent_states`, whether those are dicts of tensors (transformers 5.17) or
|
| 67 |
+
single tensors, and ignores anything not present."""
|
| 68 |
+
out = []
|
| 69 |
+
for layer in getattr(cache, "layers", None) or []:
|
| 70 |
+
for attr in STATE_NAMES:
|
| 71 |
+
v = getattr(layer, attr, None)
|
| 72 |
+
if isinstance(v, torch.Tensor):
|
| 73 |
+
if v.numel():
|
| 74 |
+
out.append((layer, attr))
|
| 75 |
+
elif isinstance(v, dict):
|
| 76 |
+
for k, t in v.items():
|
| 77 |
+
if isinstance(t, torch.Tensor) and t.numel():
|
| 78 |
+
out.append((v, k))
|
| 79 |
+
return out
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def cache_row_bytes(cache):
|
| 83 |
+
"""Bytes one row of this cache holds, summed over the layers and over `STATE_NAMES`."""
|
| 84 |
+
return sum(_get(c, k).nbytes // max(_get(c, k).shape[0], 1) for c, k in cache_state_tensors(cache))
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def unknown_state_names(cache):
|
| 88 |
+
"""Attribute names the cache's layers hold that carry tensors and are not in `STATE_NAMES`.
|
| 89 |
+
|
| 90 |
+
A tensor outside the enumerated set may or may not have a batch dimension, and `fork_cache` would leave it at one
|
| 91 |
+
row. When this is not empty the caller must not chunk."""
|
| 92 |
+
out = set()
|
| 93 |
+
for layer in getattr(cache, "layers", None) or []:
|
| 94 |
+
for name, v in vars(layer).items():
|
| 95 |
+
if name in STATE_NAMES:
|
| 96 |
+
continue
|
| 97 |
+
if isinstance(v, torch.Tensor) or (isinstance(v, dict) and any(isinstance(t, torch.Tensor) for t in v.values())):
|
| 98 |
+
out.add(name)
|
| 99 |
+
return out
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
def fork_budget_bytes(device=None, gb=None):
|
| 103 |
+
"""DECIDER_SHARED_FORK_GB, capped at half of the memory free on the device right now.
|
| 104 |
+
|
| 105 |
+
The cap applies only when the CUDA memory queries succeed: on CPU, on MPS and when the driver does not know the
|
| 106 |
+
device string, the configured budget is kept as it is."""
|
| 107 |
+
gb = float(os.environ.get("DECIDER_SHARED_FORK_GB", DEFAULT_FORK_GB)) if gb is None else float(gb)
|
| 108 |
+
budget = int(gb * (1 << 30))
|
| 109 |
+
try:
|
| 110 |
+
free, _ = torch.cuda.mem_get_info(device)
|
| 111 |
+
free += torch.cuda.memory_reserved(device) - torch.cuda.memory_allocated(device)
|
| 112 |
+
budget = min(budget, free // 2)
|
| 113 |
+
except Exception: # no CUDA device, or a device string the driver does not know
|
| 114 |
+
pass
|
| 115 |
+
return max(int(budget), 1)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def chunk_rows(prefix_bytes, n, budget_bytes=None, device=None):
|
| 119 |
+
"""Rows per fork: as many copies of the prefix cache as the budget holds, at least 1 and at most n. One row is the
|
| 120 |
+
minimum even when a single copy is over the budget: the budget bounds the fork, it cannot make it free."""
|
| 121 |
+
if budget_bytes is None:
|
| 122 |
+
budget_bytes = fork_budget_bytes(device)
|
| 123 |
+
if prefix_bytes <= 0:
|
| 124 |
+
return max(1, int(n))
|
| 125 |
+
return max(1, min(int(n), int(budget_bytes // prefix_bytes)))
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
def fork_cache(cache, m, row=0):
|
| 129 |
+
"""A new cache holding `m` copies of `cache`'s row `row`. The original is not read from again and not modified."""
|
| 130 |
+
fork = copy.copy(cache)
|
| 131 |
+
layers = getattr(cache, "layers", None)
|
| 132 |
+
if layers is not None:
|
| 133 |
+
new = []
|
| 134 |
+
for layer in layers:
|
| 135 |
+
nl = copy.copy(layer)
|
| 136 |
+
for k, v in list(vars(nl).items()): # the per-state dicts are mutated by the forward: give the fork its own
|
| 137 |
+
if isinstance(v, dict):
|
| 138 |
+
setattr(nl, k, dict(v))
|
| 139 |
+
new.append(nl)
|
| 140 |
+
fork.layers = new
|
| 141 |
+
idx = {}
|
| 142 |
+
for container, key in cache_state_tensors(fork):
|
| 143 |
+
t = _get(container, key)
|
| 144 |
+
i = idx.get(t.device)
|
| 145 |
+
if i is None:
|
| 146 |
+
i = idx[t.device] = torch.full((m,), row, dtype=torch.long, device=t.device)
|
| 147 |
+
_set(container, key, t.index_select(0, i)) # a fresh contiguous tensor: the fork never shares storage
|
| 148 |
+
return fork
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
def _count(engine, key):
|
| 152 |
+
stats = getattr(engine, "stats", None)
|
| 153 |
+
if isinstance(stats, dict):
|
| 154 |
+
stats[key] = stats.get(key, 0) + 1
|
| 155 |
+
|
| 156 |
+
|
| 157 |
+
@torch.no_grad()
|
| 158 |
+
def score_shared(engine, items, temperature=1.0, min_prefix=192, budget_bytes=None, rows_per_fork=None):
|
| 159 |
+
"""Score `items` through the shared prefix. -> one probability tensor per item, in item order, or None when the
|
| 160 |
+
request does not qualify (fewer than two rows, or a common prefix below `min_prefix`) and the caller should use
|
| 161 |
+
`score_items`.
|
| 162 |
+
|
| 163 |
+
`temperature`: a number, or one entry per item (decider.temperature.slot_temperatures).
|
| 164 |
+
`rows_per_fork` forces the chunk size; it exists for the tests that compare chunked against unchunked answers."""
|
| 165 |
+
ids = [it["ids"] for it in items]
|
| 166 |
+
n = len(ids)
|
| 167 |
+
if n < 2:
|
| 168 |
+
return None
|
| 169 |
+
lcp = common_prefix_len(ids)
|
| 170 |
+
if lcp < min_prefix:
|
| 171 |
+
return None
|
| 172 |
+
slot_temperatures(temperature, items) # a length mismatch fails before any forward
|
| 173 |
+
core, W, dev, pad = engine.core, engine.W, engine.dev, engine.tok.pad_token_id
|
| 174 |
+
pre = torch.tensor(ids[0][:lcp], device=dev)[None]
|
| 175 |
+
cache = core(input_ids=pre, use_cache=True).past_key_values
|
| 176 |
+
unknown = unknown_state_names(cache)
|
| 177 |
+
if unknown: # a state we cannot fork row by row: one fork of everything, as in 1.1.0
|
| 178 |
+
_count(engine, "shared_unchunked_layout")
|
| 179 |
+
m = n
|
| 180 |
+
elif rows_per_fork:
|
| 181 |
+
m = int(rows_per_fork)
|
| 182 |
+
else:
|
| 183 |
+
m = chunk_rows(cache_row_bytes(cache), n, budget_bytes, dev)
|
| 184 |
+
m = max(1, min(m, n))
|
| 185 |
+
out = []
|
| 186 |
+
for i in range(0, n, m):
|
| 187 |
+
part = items[i:i + m]
|
| 188 |
+
b = len(part)
|
| 189 |
+
fork = fork_cache(cache, b)
|
| 190 |
+
Ts = max(len(it["ids"]) for it in part) - lcp
|
| 191 |
+
suf = fill_ids([it["ids"][lcp:] for it in part], b, Ts, pad)
|
| 192 |
+
h = core(input_ids=suf.to(dev), past_key_values=fork, use_cache=True).last_hidden_state
|
| 193 |
+
rows = [j for j, it in enumerate(part) for _ in it["slots"]]
|
| 194 |
+
sl = [s - lcp for it in part for s in it["slots"]]
|
| 195 |
+
idx = torch.tensor([rows, sl], device=dev)
|
| 196 |
+
out += read_slots(F.linear(h[idx[0], idx[1]], W).float()[:, None, :], list(range(len(rows))), [0] * len(rows),
|
| 197 |
+
[k for it in part for k in it["nopts"]],
|
| 198 |
+
slot_temperatures(item_slice(temperature, i, i + b), part), [len(it["slots"]) for it in part])
|
| 199 |
+
del fork, h, suf # drop this chunk's fork before the next one is built
|
| 200 |
+
return out
|
decider/systemone.py
ADDED
|
@@ -0,0 +1,199 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Jev-shaped requests on top of the decider prompt format (same wire format as TypeSafe's POST /v1/systemone).
|
| 2 |
+
|
| 3 |
+
state str | dict | list JSON state is serialised compactly; questions may name a part by path (`ticket.messages[0].text`)
|
| 4 |
+
questions {id: {"type": "choice", "instructions": ..., "criteria": {name: description | {...} | [...] | None}} up to 255 options
|
| 5 |
+
{"type": "score", "instructions": ..., "criteria": [level 0 description, level 1 description, ...]} 2..10 levels
|
| 6 |
+
{"type": "noul", "instructions": ... (optional), "criteria": {"true": ..., "false": ...} (optional)}}
|
| 7 |
+
a noul question needs instructions or at least one true/false description.
|
| 8 |
+
ids are never shown to the model. `instructions` and every description may be a string or any JSON value.
|
| 9 |
+
"""
|
| 10 |
+
import json, math
|
| 11 |
+
|
| 12 |
+
MAX_CHOICE, MAX_LEVELS = 255, 10
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _txt(v):
|
| 16 |
+
return v if isinstance(v, str) else json.dumps(v, ensure_ascii=False)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
ANNOTATE_MIN = 8
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def annotate_indices(x, min_len=ANNOTATE_MIN):
|
| 23 |
+
"""Write each element's position into long arrays ({"_index": i, ...}). A path such as `records[47].text` otherwise makes
|
| 24 |
+
the model count 47 elements; with the index written down it is a lookup (json_k64 probe: 0.49 -> 0.57 accuracy)."""
|
| 25 |
+
if isinstance(x, list):
|
| 26 |
+
if len(x) >= min_len:
|
| 27 |
+
return [({"_index": i, **annotate_indices(v, min_len)} if isinstance(v, dict) else {"_index": i, "value": annotate_indices(v, min_len)}) for i, v in enumerate(x)]
|
| 28 |
+
return [annotate_indices(v, min_len) for v in x]
|
| 29 |
+
if isinstance(x, dict):
|
| 30 |
+
return {k: annotate_indices(v, min_len) for k, v in x.items()}
|
| 31 |
+
return x
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def render_state(state, index_arrays=True):
|
| 35 |
+
if isinstance(state, str):
|
| 36 |
+
return state
|
| 37 |
+
return json.dumps(annotate_indices(state) if index_arrays else state, ensure_ascii=False)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def render_question(spec):
|
| 41 |
+
"""-> dict(question=str, options=[str], type=..., names=[...]) (names: what the answer reports for each option)"""
|
| 42 |
+
t = spec.get("type", "choice"); crit = spec.get("criteria", spec.get("options"))
|
| 43 |
+
raw = spec.get("instructions", spec.get("question", ""))
|
| 44 |
+
if t in ("noul", "bool") and raw in (None, ""): # the criteria carry the question (NOUL_WITHOUT_INSTRUCTIONS)
|
| 45 |
+
ins = NOUL_WITHOUT_INSTRUCTIONS
|
| 46 |
+
described = isinstance(crit, dict) and any(d not in (None, "") for d in (crit.get("true", crit.get(True)), crit.get("false", crit.get(False))))
|
| 47 |
+
if not described and (crit is None or isinstance(crit, dict)): # a non-map is rejected below with the criteria message
|
| 48 |
+
raise ValueError("noul question without instructions: criteria must describe true or false")
|
| 49 |
+
else:
|
| 50 |
+
ins = _txt(raw)
|
| 51 |
+
if not ins:
|
| 52 |
+
raise ValueError("question without instructions")
|
| 53 |
+
if t == "choice":
|
| 54 |
+
if isinstance(crit, (list, tuple)):
|
| 55 |
+
crit = {str(c): None for c in crit}
|
| 56 |
+
if not isinstance(crit, dict) or not 2 <= len(crit) <= MAX_CHOICE:
|
| 57 |
+
raise ValueError(f"choice criteria: a map of 2..{MAX_CHOICE} options")
|
| 58 |
+
names = list(crit); opts = [n if crit[n] in (None, "") else f"{n}: {_txt(crit[n])}" for n in names]
|
| 59 |
+
elif t == "score":
|
| 60 |
+
if isinstance(crit, dict): # legend form {"0": "...", "1": "..."}
|
| 61 |
+
crit = [crit[k] for k in sorted(crit, key=float)]
|
| 62 |
+
if not isinstance(crit, (list, tuple)) or not 2 <= len(crit) <= MAX_LEVELS:
|
| 63 |
+
raise ValueError(f"score criteria: an ordered list of 2..{MAX_LEVELS} level descriptions")
|
| 64 |
+
names = list(range(len(crit))); opts = [f"{i}: {_txt(c)}" for i, c in enumerate(crit)]
|
| 65 |
+
elif t in ("noul", "bool"):
|
| 66 |
+
if crit is not None and not isinstance(crit, dict):
|
| 67 |
+
raise ValueError("noul criteria: a map of optional true/false descriptions")
|
| 68 |
+
names = [False, True]; c = crit if crit is not None else {}
|
| 69 |
+
f, tr = c.get("false", c.get(False)), c.get("true", c.get(True))
|
| 70 |
+
opts = ["no" if f in (None, "") else f"no: {_txt(f)}", "yes" if tr in (None, "") else f"yes: {_txt(tr)}"]
|
| 71 |
+
else:
|
| 72 |
+
raise ValueError(f"unknown question type {t!r}")
|
| 73 |
+
return dict(question=ins, options=opts, type="noul" if t == "bool" else t, names=names, legend=[_txt(c) for c in crit] if t == "score" else None,
|
| 74 |
+
isolated=bool(spec.get("isolated", True)))
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
# A noul question may omit `instructions` (TypeSafe's OpenAPI file marks it optional). The question id is never shown to the
|
| 78 |
+
# model, so the question text is this fixed sentence and the true/false descriptions, rendered as the options "no: ..." and
|
| 79 |
+
# "yes: ...", say what is being asked. A request that gives instructions is rendered exactly as before.
|
| 80 |
+
NOUL_WITHOUT_INSTRUCTIONS = "Which answer fits the context?"
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
# ---- isolated levels: every Score level is judged in its own row, without its number or its neighbours
|
| 84 |
+
ISOLATED = "{q}\nProposed answer: {level}\nDoes the proposed answer fit?"
|
| 85 |
+
_NUM = None
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def strip_level_number(text):
|
| 89 |
+
""""2: somewhat" -> "somewhat" (dataset legends carry the number; an isolated level must not)."""
|
| 90 |
+
import re
|
| 91 |
+
return re.sub(r"^\s*-?\d+\s*:\s*", "", text)
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
def isolated_rows(question, levels):
|
| 95 |
+
"""-> one yes/no question per level: [(question text, ["no", "yes"])]."""
|
| 96 |
+
return [(ISOLATED.format(q=question, level=strip_level_number(l)), ["no", "yes"]) for l in levels]
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def combine_isolated(p_yes):
|
| 100 |
+
"""Per-level P(fits), each computed without reference to any other level -> a distribution over levels.
|
| 101 |
+
Also returns the unnormalised mass: near 1 when exactly one level fits, low when none does, high when several do."""
|
| 102 |
+
tot = sum(p_yes) or 1e-9
|
| 103 |
+
return [x / tot for x in p_yes], tot
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def plan_rows(rqs, isolated=True):
|
| 107 |
+
"""One scoring row per question; a Score question with isolated levels becomes one yes/no row per level.
|
| 108 |
+
-> (rows [{"question", "options"}], index [(id, "iso" | "list", first row, n rows)])"""
|
| 109 |
+
rows, index = [], []
|
| 110 |
+
for k, r in rqs.items():
|
| 111 |
+
if isolated and r["type"] == "score" and r.get("isolated", True):
|
| 112 |
+
rws = isolated_rows(r["question"], r["legend"]); index.append((k, "iso", len(rows), len(rws))); rows += [dict(question=t, options=o) for t, o in rws]
|
| 113 |
+
else:
|
| 114 |
+
index.append((k, "list", len(rows), 1)); rows.append(dict(question=r["question"], options=r["options"]))
|
| 115 |
+
return rows, index
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def row_types(rqs, index):
|
| 119 |
+
"""The answer type ("choice", "noul" or "score") of every plan_rows row, in row order. An isolated Score question's
|
| 120 |
+
yes/no level rows carry "score": together they are one Score answer (decider.temperature)."""
|
| 121 |
+
types = [None] * sum(n for _, _, _, n in index)
|
| 122 |
+
for k, _, s, n in index:
|
| 123 |
+
types[s:s + n] = [rqs[k]["type"]] * n
|
| 124 |
+
return types
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def assemble(rqs, index, probs):
|
| 128 |
+
"""probs: one probability list per row (plan_rows order) -> {id: answer}."""
|
| 129 |
+
out = {}
|
| 130 |
+
for k, kind, s, n in index:
|
| 131 |
+
if kind == "iso":
|
| 132 |
+
fit = [float(probs[s + j][1]) for j in range(n)]; p, mass = combine_isolated(fit); a = format_answer(rqs[k], p)
|
| 133 |
+
a["level_fit"] = {str(j): round(x, 4) for j, x in enumerate(fit)}; a["fit_mass"] = round(mass, 4); out[k] = a
|
| 134 |
+
else:
|
| 135 |
+
out[k] = format_answer(rqs[k], probs[s])
|
| 136 |
+
return out
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def certainty(p):
|
| 140 |
+
"""1 - normalised entropy: 1 when all mass is on one option, 0 when the distribution is flat."""
|
| 141 |
+
h = -sum(x * math.log(x) for x in p if x > 0)
|
| 142 |
+
return max(0.0, 1.0 - h / math.log(len(p))) if len(p) > 1 else 1.0
|
| 143 |
+
|
| 144 |
+
|
| 145 |
+
def _clip01(x):
|
| 146 |
+
return min(1.0, max(0.0, x))
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def _normalised(p):
|
| 150 |
+
"""As the adapter's _normalize: a distribution with zero total counts as uniform."""
|
| 151 |
+
tot = sum(p)
|
| 152 |
+
return [1.0 / len(p)] * len(p) if tot == 0 else [x / tot for x in p]
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def choice_confidence(p):
|
| 156 |
+
"""TypeSafe's Choice confidence: the largest probability rescaled so that a uniform distribution gives 0 and all mass on one
|
| 157 |
+
option gives 1, (n * p_max - 1) / (n - 1); 1 for a single option."""
|
| 158 |
+
n = len(p); p = _normalised(p)
|
| 159 |
+
return 1.0 if n <= 1 else _clip01((n * max(p) - 1) / (n - 1))
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def score_confidence(p):
|
| 163 |
+
"""TypeSafe's Score confidence (system-one-adapter-python, confidence_metrics.score_confidence): 1 minus the expected distance
|
| 164 |
+
from the most likely level, divided by D = mean over levels i of |i - (n - 1)/2| (the mean distance of the levels from the
|
| 165 |
+
middle of the scale), floored at 0; 1 for a single level.
|
| 166 |
+
With two levels it equals choice_confidence."""
|
| 167 |
+
n = len(p); p = _normalised(p)
|
| 168 |
+
if n <= 1:
|
| 169 |
+
return 1.0
|
| 170 |
+
k = max(range(n), key=p.__getitem__)
|
| 171 |
+
spread = sum(x * abs(i - k) for i, x in enumerate(p))
|
| 172 |
+
uniform = sum(abs(i - (n - 1) / 2) for i in range(n)) / n
|
| 173 |
+
return _clip01(1.0 - spread / uniform)
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
def format_answer(rq, p, nd=4):
|
| 177 |
+
"""rq: render_question output; p: probabilities in option order.
|
| 178 |
+
`confidence` is TypeSafe's (choice_confidence / score_confidence); `x_p_max` is the largest probability, which was
|
| 179 |
+
`confidence` before 1.3.0."""
|
| 180 |
+
p = [float(x) for x in p[:len(rq["options"])]]; s = sum(p) or 1.0; p = [x / s for x in p]
|
| 181 |
+
j = max(range(len(p)), key=p.__getitem__)
|
| 182 |
+
if rq["type"] == "noul":
|
| 183 |
+
return {"type": "noul", "noul": round(p[1], nd)}
|
| 184 |
+
if rq["type"] == "choice":
|
| 185 |
+
return {"type": "choice", "choice": rq["names"][j], "confidence": round(choice_confidence(p), nd), "x_p_max": round(p[j], nd),
|
| 186 |
+
"certainty": round(certainty(p), nd), "probabilities": {n: round(x, nd) for n, x in zip(rq["names"], p)}}
|
| 187 |
+
return {"type": "score", "score": round(sum(i * x for i, x in enumerate(p)), 2), "confidence": round(score_confidence(p), nd), "x_p_max": round(p[j], nd),
|
| 188 |
+
"certainty": round(certainty(p), nd),
|
| 189 |
+
"legend": {str(i): d for i, d in enumerate(rq["legend"])}, "probabilities": {str(i): round(x, nd) for i, x in enumerate(p)}}
|
| 190 |
+
|
| 191 |
+
|
| 192 |
+
def unique_tokens(items):
|
| 193 |
+
"""Input tokens of a request whose rows share a prefix (the state): the prefix counts once."""
|
| 194 |
+
ids = [it["ids"] for it in items]
|
| 195 |
+
if len(ids) < 2:
|
| 196 |
+
return sum(len(x) for x in ids)
|
| 197 |
+
lcp = 0; short = min(len(x) for x in ids)
|
| 198 |
+
while lcp < short and all(x[lcp] == ids[0][lcp] for x in ids): lcp += 1
|
| 199 |
+
return lcp + sum(len(x) - lcp for x in ids)
|
decider/temperature.py
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Answer temperatures from decider_config.json (1.4.0).
|
| 2 |
+
|
| 3 |
+
Every answer is softmax(logits / T) over its option letters. decider_config.json sets T:
|
| 4 |
+
|
| 5 |
+
"temperature": 1.3 one value for every answer (the only form before 1.4.0)
|
| 6 |
+
"temperature_by_type": {"choice": 1.48, "noul": 2.22, "score": 1.38}
|
| 7 |
+
optional; one value per answer type, a missing type uses "temperature"
|
| 8 |
+
"temperature_schema_first": 1.18 optional; the schema cache (questions-first layout), as before
|
| 9 |
+
"temperature_schema_first_by_type": {...} optional; per answer type on the schema cache
|
| 10 |
+
|
| 11 |
+
The answer types are the /v1/systemone question types. A /decide field maps onto them: "choice" -> "choice", "bool" -> "noul",
|
| 12 |
+
"scale" -> "score". Plain `Decider.decide` questions (a question and its options, no type) are "choice". A Score question
|
| 13 |
+
read with isolated levels (one yes/no row per level) uses the "score" temperature on every one of its level rows: the rows form
|
| 14 |
+
one Score answer, and a temperature fitted on Score answers is fitted through that same readout (decider.calibrate).
|
| 15 |
+
|
| 16 |
+
On the state-first layout the temperature of an answer of type t is temperature_by_type[t], else temperature. On the schema
|
| 17 |
+
cache it is temperature_schema_first_by_type[t], else temperature_schema_first, else (no schema-first value at all) the
|
| 18 |
+
state-first temperature of t. An explicit override (Decider(temperature=...), DECIDER_TEMPERATURE) replaces the state-first
|
| 19 |
+
"temperature" and switches temperature_by_type off, so the override is the one temperature of every state-first answer, as it
|
| 20 |
+
was before 1.4.0.
|
| 21 |
+
|
| 22 |
+
A config without a by-type map gives every path the single number it gave in 1.3.0, and that number reaches the engines as the
|
| 23 |
+
same Python float, so the probabilities are bit-identical to 1.3.0.
|
| 24 |
+
"""
|
| 25 |
+
import math
|
| 26 |
+
|
| 27 |
+
TYPES = ("choice", "noul", "score")
|
| 28 |
+
FIELD_TYPES = {"choice": "choice", "bool": "noul", "scale": "score"} # /decide field type -> answer type
|
| 29 |
+
_KEYS_TEXT = ('the keys are "choice", "noul" and "score" (a /decide "bool" field is "noul", a "scale" field is "score"; '
|
| 30 |
+
'a /v1/systemone "bool" question is "noul")')
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def positive(value, where):
|
| 34 |
+
"""A temperature: a number (or a numeric string, as float() reads it, which 1.3.0 accepted for "temperature") that is finite
|
| 35 |
+
and > 0. Raises ValueError naming `where`."""
|
| 36 |
+
if isinstance(value, bool):
|
| 37 |
+
raise ValueError(f"{where} must be a finite number > 0, got {value!r}")
|
| 38 |
+
try:
|
| 39 |
+
v = float(value)
|
| 40 |
+
except (TypeError, ValueError):
|
| 41 |
+
raise ValueError(f"{where} must be a finite number > 0, got {value!r}") from None
|
| 42 |
+
if not math.isfinite(v) or v <= 0:
|
| 43 |
+
raise ValueError(f"{where} must be a finite number > 0, got {value!r}")
|
| 44 |
+
return v
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def by_type(m, where):
|
| 48 |
+
"""Validate a {answer type: temperature} map. None -> {}. Unknown keys, non-numbers and values that are not finite and > 0
|
| 49 |
+
raise ValueError."""
|
| 50 |
+
if m is None:
|
| 51 |
+
return {}
|
| 52 |
+
if not isinstance(m, dict):
|
| 53 |
+
raise ValueError(f"{where} must be a map {{answer type: temperature}}, got {type(m).__name__}; " + _KEYS_TEXT)
|
| 54 |
+
out = {}
|
| 55 |
+
for k, v in m.items():
|
| 56 |
+
if k not in TYPES:
|
| 57 |
+
raise ValueError(f"{where} has the unknown key {k!r}; " + _KEYS_TEXT)
|
| 58 |
+
if isinstance(v, bool) or not isinstance(v, (int, float)):
|
| 59 |
+
raise ValueError(f"{where}[{k!r}] must be a finite number > 0, got {v!r}")
|
| 60 |
+
out[k] = positive(v, f"{where}[{k!r}]")
|
| 61 |
+
return out
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def from_config(cfg, temperature=None, temperature_by_type=None):
|
| 65 |
+
"""-> ((T, by_type) for the state-first layout, (T, by_type) for the schema cache).
|
| 66 |
+
|
| 67 |
+
temperature / temperature_by_type: explicit overrides (Decider arguments, DECIDER_TEMPERATURE). An explicit temperature
|
| 68 |
+
without an explicit map switches the config's map off (see the module docstring)."""
|
| 69 |
+
cfg = cfg or {}
|
| 70 |
+
where = "decider_config.json"
|
| 71 |
+
if temperature is not None:
|
| 72 |
+
T = positive(temperature, "temperature")
|
| 73 |
+
m = by_type(temperature_by_type, "temperature_by_type") if temperature_by_type is not None else {}
|
| 74 |
+
else:
|
| 75 |
+
T = positive(cfg.get("temperature", 1.0), f'{where} "temperature"')
|
| 76 |
+
m = (by_type(temperature_by_type, "temperature_by_type") if temperature_by_type is not None
|
| 77 |
+
else by_type(cfg.get("temperature_by_type"), f'{where} "temperature_by_type"'))
|
| 78 |
+
ms = by_type(cfg.get("temperature_schema_first_by_type"), f'{where} "temperature_schema_first_by_type"')
|
| 79 |
+
if "temperature_schema_first" in cfg:
|
| 80 |
+
schema = (positive(cfg["temperature_schema_first"], f'{where} "temperature_schema_first"'), ms)
|
| 81 |
+
else:
|
| 82 |
+
schema = (T, {**m, **ms})
|
| 83 |
+
return (T, m), schema
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def effective(T, m):
|
| 87 |
+
"""{answer type: the temperature it gets} for reporting (/health, the ready line)."""
|
| 88 |
+
return {t: m.get(t, T) for t in TYPES}
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def for_types(T, m, types):
|
| 92 |
+
"""One temperature per slot, or the scalar T itself when there is no map (the 1.3.0 call)."""
|
| 93 |
+
if not m:
|
| 94 |
+
return T
|
| 95 |
+
return [m.get(t, T) for t in types]
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
def item_types(it):
|
| 99 |
+
"""The answer type of every slot of a prompt item ("types", set where the item is built); None for an item without it."""
|
| 100 |
+
ts = it.get("types")
|
| 101 |
+
return list(ts) if ts is not None else [None] * len(it["slots"])
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
def for_items(T, m, items):
|
| 105 |
+
"""The `temperature` argument of Engine.score_items / score_shared: the scalar T when there is no map (the 1.3.0 call),
|
| 106 |
+
else one list of per-slot temperatures per item."""
|
| 107 |
+
if not m:
|
| 108 |
+
return T
|
| 109 |
+
return [[m.get(t, T) for t in item_types(it)] for it in items]
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
def slot_temperatures(temperature, items):
|
| 113 |
+
"""Engine side. temperature: a scalar (returned as it is), or one entry per item, each a scalar or one value per slot of
|
| 114 |
+
that item. -> the scalar, or a flat list with one temperature per slot in item order."""
|
| 115 |
+
if not isinstance(temperature, (list, tuple)):
|
| 116 |
+
return temperature
|
| 117 |
+
if len(temperature) != len(items):
|
| 118 |
+
raise ValueError(f"temperature: {len(temperature)} entries for {len(items)} items")
|
| 119 |
+
flat = []
|
| 120 |
+
for t, it in zip(temperature, items):
|
| 121 |
+
n = len(it["slots"])
|
| 122 |
+
if isinstance(t, (list, tuple)):
|
| 123 |
+
if len(t) != n:
|
| 124 |
+
raise ValueError(f"temperature: {len(t)} values for an item with {n} slots")
|
| 125 |
+
flat += list(t)
|
| 126 |
+
else:
|
| 127 |
+
flat += [t] * n
|
| 128 |
+
return flat
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def item_slice(temperature, lo, hi):
|
| 132 |
+
"""The per-item temperature entries of items[lo:hi] (a scalar is shared by every item)."""
|
| 133 |
+
return temperature[lo:hi] if isinstance(temperature, (list, tuple)) else temperature
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def scaled_softmax(lg, temperature):
|
| 137 |
+
"""softmax(lg / T) over the last axis. A scalar T is the 1.3.0 expression unchanged; a list gives one T per row of lg."""
|
| 138 |
+
import torch
|
| 139 |
+
if isinstance(temperature, (list, tuple)):
|
| 140 |
+
if len(temperature) != lg.shape[0]:
|
| 141 |
+
raise ValueError(f"temperature: {len(temperature)} values for {lg.shape[0]} slots")
|
| 142 |
+
t = torch.tensor(temperature, dtype=lg.dtype).to(lg.device, non_blocking=True)[:, None]
|
| 143 |
+
return torch.softmax(lg / t, -1)
|
| 144 |
+
return torch.softmax(lg / temperature, -1)
|
decider_config.json
ADDED
|
@@ -0,0 +1,18 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"temperature": 1.097,
|
| 3 |
+
"temperature_by_type": {
|
| 4 |
+
"noul": 1.097
|
| 5 |
+
},
|
| 6 |
+
"neutralize_none": false,
|
| 7 |
+
"version": "4b-v2.1-tr-lora-tr-lora",
|
| 8 |
+
"base": "hayriyigit/decider-4b-tr@df9aacb9f4c9e4e64178176502f51ae4db271f26 + LoRA (merged)",
|
| 9 |
+
"layout": "plain",
|
| 10 |
+
"max_options": 255,
|
| 11 |
+
"max_state_tokens": 32768,
|
| 12 |
+
"schema_first": false,
|
| 13 |
+
"schema_first_trained": false,
|
| 14 |
+
"isolated_levels": true,
|
| 15 |
+
"release_date": "2026-09-24",
|
| 16 |
+
"requires": "decider-ai>=1.4.0 for temperature_by_type; older versions serve every answer at temperature",
|
| 17 |
+
"stage": "LoRA rank 64 (alpha 128) on q_proj, k_proj, v_proj, o_proj, in_proj_qkv, in_proj_z, out_proj, gate_proj, up_proj, down_proj; LR 0.0001, 1 epochs; cross-entropy on Turkish decision rows, KL toward the base model on English replay rows; temperatures fitted with decider.calibrate.fit_by_type on held-out Turkish rows"
|
| 18 |
+
}
|
eval_report.json
ADDED
|
@@ -0,0 +1,331 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"model": "decider-4b-v2.1-tr-lora-tr-lora",
|
| 3 |
+
"path": "/workspace/laya/answerability_run/model",
|
| 4 |
+
"revision": "eb5fbdfc9448473ec25e399882912863afbdb70e",
|
| 5 |
+
"fp8_layers": 0,
|
| 6 |
+
"limit": 0,
|
| 7 |
+
"skipped": {},
|
| 8 |
+
"seconds": 502,
|
| 9 |
+
"report": {
|
| 10 |
+
"all": {
|
| 11 |
+
"n": 16243,
|
| 12 |
+
"accuracy": 0.8931,
|
| 13 |
+
"brier": 0.1653,
|
| 14 |
+
"nll": 0.2898,
|
| 15 |
+
"ece": 0.0426
|
| 16 |
+
},
|
| 17 |
+
"source/hotpot_tr": {
|
| 18 |
+
"n": 1500,
|
| 19 |
+
"accuracy": 0.98,
|
| 20 |
+
"brier": 0.032,
|
| 21 |
+
"nll": 0.0605,
|
| 22 |
+
"ece": 0.0063
|
| 23 |
+
},
|
| 24 |
+
"source/hotpot_tr/noul": {
|
| 25 |
+
"n": 1500,
|
| 26 |
+
"accuracy": 0.98,
|
| 27 |
+
"brier": 0.032,
|
| 28 |
+
"nll": 0.0605,
|
| 29 |
+
"ece": 0.0063
|
| 30 |
+
},
|
| 31 |
+
"source/musique_tr": {
|
| 32 |
+
"n": 500,
|
| 33 |
+
"accuracy": 0.898,
|
| 34 |
+
"brier": 0.1518,
|
| 35 |
+
"nll": 0.2529,
|
| 36 |
+
"ece": 0.0296
|
| 37 |
+
},
|
| 38 |
+
"source/musique_tr/noul": {
|
| 39 |
+
"n": 500,
|
| 40 |
+
"accuracy": 0.898,
|
| 41 |
+
"brier": 0.1518,
|
| 42 |
+
"nll": 0.2529,
|
| 43 |
+
"ece": 0.0296
|
| 44 |
+
},
|
| 45 |
+
"source/squad_tr": {
|
| 46 |
+
"n": 11873,
|
| 47 |
+
"accuracy": 0.8614,
|
| 48 |
+
"brier": 0.2146,
|
| 49 |
+
"nll": 0.3759,
|
| 50 |
+
"ece": 0.0572
|
| 51 |
+
},
|
| 52 |
+
"source/squad_tr/noul": {
|
| 53 |
+
"n": 11873,
|
| 54 |
+
"accuracy": 0.8614,
|
| 55 |
+
"brier": 0.2146,
|
| 56 |
+
"nll": 0.3759,
|
| 57 |
+
"ece": 0.0572
|
| 58 |
+
},
|
| 59 |
+
"source/synth_tr": {
|
| 60 |
+
"n": 1170,
|
| 61 |
+
"accuracy": 0.9923,
|
| 62 |
+
"brier": 0.0113,
|
| 63 |
+
"nll": 0.023,
|
| 64 |
+
"ece": 0.0038
|
| 65 |
+
},
|
| 66 |
+
"source/synth_tr/noul": {
|
| 67 |
+
"n": 1170,
|
| 68 |
+
"accuracy": 0.9923,
|
| 69 |
+
"brier": 0.0113,
|
| 70 |
+
"nll": 0.023,
|
| 71 |
+
"ece": 0.0038
|
| 72 |
+
},
|
| 73 |
+
"source/wiki2_tr": {
|
| 74 |
+
"n": 1200,
|
| 75 |
+
"accuracy": 1.0,
|
| 76 |
+
"brier": 0.0001,
|
| 77 |
+
"nll": 0.0006,
|
| 78 |
+
"ece": 0.0006
|
| 79 |
+
},
|
| 80 |
+
"source/wiki2_tr/noul": {
|
| 81 |
+
"n": 1200,
|
| 82 |
+
"accuracy": 1.0,
|
| 83 |
+
"brier": 0.0001,
|
| 84 |
+
"nll": 0.0006,
|
| 85 |
+
"ece": 0.0006
|
| 86 |
+
},
|
| 87 |
+
"type/noul": {
|
| 88 |
+
"n": 16243,
|
| 89 |
+
"accuracy": 0.8931,
|
| 90 |
+
"brier": 0.1653,
|
| 91 |
+
"nll": 0.2898,
|
| 92 |
+
"ece": 0.0426
|
| 93 |
+
},
|
| 94 |
+
"workflow/hotpot_tr/all_hops": {
|
| 95 |
+
"n": 500,
|
| 96 |
+
"accuracy": 0.97,
|
| 97 |
+
"brier": 0.0527,
|
| 98 |
+
"nll": 0.1011,
|
| 99 |
+
"ece": 0.0126
|
| 100 |
+
},
|
| 101 |
+
"workflow/hotpot_tr/distractors_only": {
|
| 102 |
+
"n": 500,
|
| 103 |
+
"accuracy": 0.996,
|
| 104 |
+
"brier": 0.0059,
|
| 105 |
+
"nll": 0.01,
|
| 106 |
+
"ece": 0.0065
|
| 107 |
+
},
|
| 108 |
+
"workflow/hotpot_tr/missing_hop": {
|
| 109 |
+
"n": 500,
|
| 110 |
+
"accuracy": 0.974,
|
| 111 |
+
"brier": 0.0374,
|
| 112 |
+
"nll": 0.0704,
|
| 113 |
+
"ece": 0.0127
|
| 114 |
+
},
|
| 115 |
+
"workflow/musique_tr/all_hops": {
|
| 116 |
+
"n": 250,
|
| 117 |
+
"accuracy": 0.824,
|
| 118 |
+
"brier": 0.2596,
|
| 119 |
+
"nll": 0.4257,
|
| 120 |
+
"ece": 0.0587
|
| 121 |
+
},
|
| 122 |
+
"workflow/musique_tr/missing_hop": {
|
| 123 |
+
"n": 250,
|
| 124 |
+
"accuracy": 0.972,
|
| 125 |
+
"brier": 0.044,
|
| 126 |
+
"nll": 0.0801,
|
| 127 |
+
"ece": 0.0199
|
| 128 |
+
},
|
| 129 |
+
"workflow/squad_tr/answerability": {
|
| 130 |
+
"n": 11873,
|
| 131 |
+
"accuracy": 0.8614,
|
| 132 |
+
"brier": 0.2146,
|
| 133 |
+
"nll": 0.3759,
|
| 134 |
+
"ece": 0.0572
|
| 135 |
+
},
|
| 136 |
+
"workflow/synth_tr/synth_answerable": {
|
| 137 |
+
"n": 352,
|
| 138 |
+
"accuracy": 0.983,
|
| 139 |
+
"brier": 0.0226,
|
| 140 |
+
"nll": 0.0439,
|
| 141 |
+
"ece": 0.0118
|
| 142 |
+
},
|
| 143 |
+
"workflow/synth_tr/synth_eksik_oznitelik": {
|
| 144 |
+
"n": 127,
|
| 145 |
+
"accuracy": 0.9921,
|
| 146 |
+
"brier": 0.0208,
|
| 147 |
+
"nll": 0.0531,
|
| 148 |
+
"ece": 0.014
|
| 149 |
+
},
|
| 150 |
+
"workflow/synth_tr/synth_kismi": {
|
| 151 |
+
"n": 232,
|
| 152 |
+
"accuracy": 0.9957,
|
| 153 |
+
"brier": 0.0055,
|
| 154 |
+
"nll": 0.0096,
|
| 155 |
+
"ece": 0.0073
|
| 156 |
+
},
|
| 157 |
+
"workflow/synth_tr/synth_kurum": {
|
| 158 |
+
"n": 113,
|
| 159 |
+
"accuracy": 1.0,
|
| 160 |
+
"brier": 0.0,
|
| 161 |
+
"nll": 0.0005,
|
| 162 |
+
"ece": 0.0005
|
| 163 |
+
},
|
| 164 |
+
"workflow/synth_tr/synth_kısmi": {
|
| 165 |
+
"n": 2,
|
| 166 |
+
"accuracy": 1.0,
|
| 167 |
+
"brier": 0.0,
|
| 168 |
+
"nll": 0.0001,
|
| 169 |
+
"ece": 0.0001
|
| 170 |
+
},
|
| 171 |
+
"workflow/synth_tr/synth_varlik": {
|
| 172 |
+
"n": 115,
|
| 173 |
+
"accuracy": 0.9913,
|
| 174 |
+
"brier": 0.0115,
|
| 175 |
+
"nll": 0.016,
|
| 176 |
+
"ece": 0.0056
|
| 177 |
+
},
|
| 178 |
+
"workflow/synth_tr/synth_varlık": {
|
| 179 |
+
"n": 3,
|
| 180 |
+
"accuracy": 1.0,
|
| 181 |
+
"brier": 0.0,
|
| 182 |
+
"nll": 0.0002,
|
| 183 |
+
"ece": 0.0002
|
| 184 |
+
},
|
| 185 |
+
"workflow/synth_tr/synth_yer_facet": {
|
| 186 |
+
"n": 113,
|
| 187 |
+
"accuracy": 1.0,
|
| 188 |
+
"brier": 0.0002,
|
| 189 |
+
"nll": 0.0025,
|
| 190 |
+
"ece": 0.0025
|
| 191 |
+
},
|
| 192 |
+
"workflow/synth_tr/synth_zaman": {
|
| 193 |
+
"n": 113,
|
| 194 |
+
"accuracy": 1.0,
|
| 195 |
+
"brier": 0.0001,
|
| 196 |
+
"nll": 0.0023,
|
| 197 |
+
"ece": 0.0023
|
| 198 |
+
},
|
| 199 |
+
"workflow/wiki2_tr/other_entity_same_attribute": {
|
| 200 |
+
"n": 600,
|
| 201 |
+
"accuracy": 1.0,
|
| 202 |
+
"brier": 0.0,
|
| 203 |
+
"nll": 0.0002,
|
| 204 |
+
"ece": 0.0002
|
| 205 |
+
},
|
| 206 |
+
"workflow/wiki2_tr/same_entity": {
|
| 207 |
+
"n": 600,
|
| 208 |
+
"accuracy": 1.0,
|
| 209 |
+
"brier": 0.0003,
|
| 210 |
+
"nll": 0.001,
|
| 211 |
+
"ece": 0.0009
|
| 212 |
+
},
|
| 213 |
+
"hotpot_tr/answerability": {
|
| 214 |
+
"n": 1500,
|
| 215 |
+
"acc_answerable": 0.97,
|
| 216 |
+
"acc_unanswerable": 0.984,
|
| 217 |
+
"auroc": 0.997
|
| 218 |
+
},
|
| 219 |
+
"hotpot_tr/all_hops/answerability": {
|
| 220 |
+
"n": 500,
|
| 221 |
+
"acc_answerable": 0.97,
|
| 222 |
+
"acc_unanswerable": null
|
| 223 |
+
},
|
| 224 |
+
"hotpot_tr/distractors_only/answerability": {
|
| 225 |
+
"n": 500,
|
| 226 |
+
"acc_answerable": null,
|
| 227 |
+
"acc_unanswerable": 0.996
|
| 228 |
+
},
|
| 229 |
+
"hotpot_tr/missing_hop/answerability": {
|
| 230 |
+
"n": 500,
|
| 231 |
+
"acc_answerable": null,
|
| 232 |
+
"acc_unanswerable": 0.972
|
| 233 |
+
},
|
| 234 |
+
"musique_tr/answerability": {
|
| 235 |
+
"n": 500,
|
| 236 |
+
"acc_answerable": 0.832,
|
| 237 |
+
"acc_unanswerable": 0.972,
|
| 238 |
+
"auroc": 0.9783
|
| 239 |
+
},
|
| 240 |
+
"musique_tr/all_hops/answerability": {
|
| 241 |
+
"n": 250,
|
| 242 |
+
"acc_answerable": 0.832,
|
| 243 |
+
"acc_unanswerable": null
|
| 244 |
+
},
|
| 245 |
+
"musique_tr/missing_hop/answerability": {
|
| 246 |
+
"n": 250,
|
| 247 |
+
"acc_answerable": null,
|
| 248 |
+
"acc_unanswerable": 0.972
|
| 249 |
+
},
|
| 250 |
+
"squad_tr/answerability": {
|
| 251 |
+
"n": 11873,
|
| 252 |
+
"acc_answerable": 0.9239,
|
| 253 |
+
"acc_unanswerable": 0.7978,
|
| 254 |
+
"auroc": 0.9346
|
| 255 |
+
},
|
| 256 |
+
"squad_tr/answerability/answerability": {
|
| 257 |
+
"n": 11873,
|
| 258 |
+
"acc_answerable": 0.9239,
|
| 259 |
+
"acc_unanswerable": 0.7978,
|
| 260 |
+
"auroc": 0.9346
|
| 261 |
+
},
|
| 262 |
+
"synth_tr/answerability": {
|
| 263 |
+
"n": 1170,
|
| 264 |
+
"acc_answerable": 0.983,
|
| 265 |
+
"acc_unanswerable": 0.9963,
|
| 266 |
+
"auroc": 0.9996
|
| 267 |
+
},
|
| 268 |
+
"synth_tr/synth_answerable/answerability": {
|
| 269 |
+
"n": 352,
|
| 270 |
+
"acc_answerable": 0.983,
|
| 271 |
+
"acc_unanswerable": null
|
| 272 |
+
},
|
| 273 |
+
"synth_tr/synth_eksik_oznitelik/answerability": {
|
| 274 |
+
"n": 127,
|
| 275 |
+
"acc_answerable": null,
|
| 276 |
+
"acc_unanswerable": 0.9921
|
| 277 |
+
},
|
| 278 |
+
"synth_tr/synth_kismi/answerability": {
|
| 279 |
+
"n": 232,
|
| 280 |
+
"acc_answerable": null,
|
| 281 |
+
"acc_unanswerable": 0.9957
|
| 282 |
+
},
|
| 283 |
+
"synth_tr/synth_kurum/answerability": {
|
| 284 |
+
"n": 113,
|
| 285 |
+
"acc_answerable": null,
|
| 286 |
+
"acc_unanswerable": 1.0
|
| 287 |
+
},
|
| 288 |
+
"synth_tr/synth_kısmi/answerability": {
|
| 289 |
+
"n": 2,
|
| 290 |
+
"acc_answerable": null,
|
| 291 |
+
"acc_unanswerable": 1.0
|
| 292 |
+
},
|
| 293 |
+
"synth_tr/synth_varlik/answerability": {
|
| 294 |
+
"n": 115,
|
| 295 |
+
"acc_answerable": null,
|
| 296 |
+
"acc_unanswerable": 0.9913
|
| 297 |
+
},
|
| 298 |
+
"synth_tr/synth_varlık/answerability": {
|
| 299 |
+
"n": 3,
|
| 300 |
+
"acc_answerable": null,
|
| 301 |
+
"acc_unanswerable": 1.0
|
| 302 |
+
},
|
| 303 |
+
"synth_tr/synth_yer_facet/answerability": {
|
| 304 |
+
"n": 113,
|
| 305 |
+
"acc_answerable": null,
|
| 306 |
+
"acc_unanswerable": 1.0
|
| 307 |
+
},
|
| 308 |
+
"synth_tr/synth_zaman/answerability": {
|
| 309 |
+
"n": 113,
|
| 310 |
+
"acc_answerable": null,
|
| 311 |
+
"acc_unanswerable": 1.0
|
| 312 |
+
},
|
| 313 |
+
"wiki2_tr/answerability": {
|
| 314 |
+
"n": 1200,
|
| 315 |
+
"acc_answerable": 1.0,
|
| 316 |
+
"acc_unanswerable": 1.0,
|
| 317 |
+
"auroc": 1.0
|
| 318 |
+
},
|
| 319 |
+
"wiki2_tr/other_entity_same_attribute/answerability": {
|
| 320 |
+
"n": 600,
|
| 321 |
+
"acc_answerable": null,
|
| 322 |
+
"acc_unanswerable": 1.0
|
| 323 |
+
},
|
| 324 |
+
"wiki2_tr/same_entity/answerability": {
|
| 325 |
+
"n": 600,
|
| 326 |
+
"acc_answerable": 1.0,
|
| 327 |
+
"acc_unanswerable": null
|
| 328 |
+
}
|
| 329 |
+
},
|
| 330 |
+
"train_key": "dc5d693207f19271"
|
| 331 |
+
}
|
export.json
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"train_key": "dc5d693207f19271",
|
| 3 |
+
"calibration": {
|
| 4 |
+
"temperature_by_type": {
|
| 5 |
+
"noul": 1.097
|
| 6 |
+
},
|
| 7 |
+
"temperature": 1.097,
|
| 8 |
+
"rows": {
|
| 9 |
+
"choice": 0,
|
| 10 |
+
"noul": 1500,
|
| 11 |
+
"score": 0
|
| 12 |
+
},
|
| 13 |
+
"nll": {
|
| 14 |
+
"noul": {
|
| 15 |
+
"T=1": 0.056,
|
| 16 |
+
"fitted": 0.0556
|
| 17 |
+
}
|
| 18 |
+
}
|
| 19 |
+
}
|
| 20 |
+
}
|
generation_config.json
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"_from_model_config": true,
|
| 3 |
+
"eos_token_id": 248044,
|
| 4 |
+
"transformers_version": "5.17.0",
|
| 5 |
+
"use_cache": true
|
| 6 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:e2d6fc070f1b6a8f967388e035f1051c5789e2e1c45e871da81fa2aac20ef614
|
| 3 |
+
size 8411558400
|
tokenizer.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:06b9509352d2af50381ab2247e083b80d32d5c0aba91c272ca9ff729b6a0e523
|
| 3 |
+
size 19989325
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"add_prefix_space": false,
|
| 3 |
+
"audio_bos_token": "<|audio_start|>",
|
| 4 |
+
"audio_eos_token": "<|audio_end|>",
|
| 5 |
+
"audio_token": "<|audio_pad|>",
|
| 6 |
+
"backend": "tokenizers",
|
| 7 |
+
"bos_token": null,
|
| 8 |
+
"clean_up_tokenization_spaces": false,
|
| 9 |
+
"eos_token": "<|endoftext|>",
|
| 10 |
+
"errors": "replace",
|
| 11 |
+
"image_token": "<|image_pad|>",
|
| 12 |
+
"is_local": true,
|
| 13 |
+
"local_files_only": false,
|
| 14 |
+
"model_max_length": 262144,
|
| 15 |
+
"model_specific_special_tokens": {
|
| 16 |
+
"audio_bos_token": "<|audio_start|>",
|
| 17 |
+
"audio_eos_token": "<|audio_end|>",
|
| 18 |
+
"audio_token": "<|audio_pad|>",
|
| 19 |
+
"image_token": "<|image_pad|>",
|
| 20 |
+
"video_token": "<|video_pad|>",
|
| 21 |
+
"vision_bos_token": "<|vision_start|>",
|
| 22 |
+
"vision_eos_token": "<|vision_end|>"
|
| 23 |
+
},
|
| 24 |
+
"pad_token": "<|endoftext|>",
|
| 25 |
+
"pretokenize_regex": "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
|
| 26 |
+
"split_special_tokens": false,
|
| 27 |
+
"tokenizer_class": "Qwen2Tokenizer",
|
| 28 |
+
"unk_token": null,
|
| 29 |
+
"video_token": "<|video_pad|>",
|
| 30 |
+
"vision_bos_token": "<|vision_start|>",
|
| 31 |
+
"vision_eos_token": "<|vision_end|>"
|
| 32 |
+
}
|