mirror of
https://github.com/open-jarvis/OpenJarvis.git
synced 2026-07-27 21:05:34 +00:00
feat(distillation): M1→M2→M3 spec-level distillation pipeline + hill-climb optimizer (#273)
This commit is contained in:
+13
@@ -98,3 +98,16 @@ src/openjarvis/channels/whatsapp_baileys_bridge/node_modules/
|
||||
# Second Repos
|
||||
Inline/
|
||||
scratch/
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Distillation runtime artifacts (defense in depth — these should always live
|
||||
# in ~/.openjarvis/, never inside the source tree, but we ignore them here in
|
||||
# case OPENJARVIS_HOME is misconfigured during dev)
|
||||
# ---------------------------------------------------------------------------
|
||||
.openjarvis/
|
||||
learning.db
|
||||
**/learning/sessions/
|
||||
**/learning/pending_review/
|
||||
**/learning/benchmarks/
|
||||
**/teacher_traces/
|
||||
*.session.json
|
||||
|
||||
Submodule
+1
Submodule Inline added at 03673aaa42
@@ -0,0 +1,450 @@
|
||||
<!DOCTYPE html>
|
||||
<html>
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<style>
|
||||
*{margin:0;padding:0;box-sizing:border-box}
|
||||
html,body,:root{
|
||||
background:transparent !important;
|
||||
background-color:transparent !important;
|
||||
font-family:-apple-system,BlinkMacSystemFont,"SF Pro Text",sans-serif;
|
||||
color:#fff;height:100%;overflow:hidden;
|
||||
}
|
||||
#container{
|
||||
display:flex;flex-direction:column;
|
||||
height:100%;padding:8px 12px;
|
||||
}
|
||||
#messages{
|
||||
flex:1;overflow-y:auto;
|
||||
display:flex;flex-direction:column;gap:6px;
|
||||
margin-bottom:8px;padding:12px;
|
||||
border-radius:16px;
|
||||
background:rgba(30,30,30,0.88);
|
||||
border:1px solid rgba(255,255,255,0.20);
|
||||
}
|
||||
#messages:empty{display:none}
|
||||
#messages::-webkit-scrollbar{width:6px}
|
||||
#messages::-webkit-scrollbar-thumb{background:rgba(255,255,255,0.15);border-radius:3px}
|
||||
#messages::-webkit-scrollbar-track{background:transparent}
|
||||
.msg{
|
||||
padding:8px 12px;border-radius:12px;
|
||||
font-size:13px;line-height:1.55;max-width:90%;
|
||||
white-space:pre-wrap;word-wrap:break-word;
|
||||
-webkit-user-select:text;user-select:text;
|
||||
}
|
||||
.msg.user{
|
||||
align-self:flex-end;
|
||||
background:rgba(59,130,246,0.35);
|
||||
color:rgba(255,255,255,0.95);
|
||||
}
|
||||
.msg.assistant{
|
||||
align-self:flex-start;
|
||||
background:rgba(255,255,255,0.20);
|
||||
color:rgba(255,255,255,0.88);
|
||||
white-space:normal;
|
||||
}
|
||||
.msg.assistant > p{margin:0 0 6px 0;white-space:pre-wrap}
|
||||
.msg.assistant > p:last-child{margin-bottom:0}
|
||||
.msg.assistant h1,.msg.assistant h2,.msg.assistant h3{
|
||||
font-size:14px;font-weight:600;margin:6px 0 4px 0;
|
||||
}
|
||||
.msg.assistant ul,.msg.assistant ol{margin:2px 0 6px 18px;padding:0}
|
||||
.msg.assistant li{margin:1px 0}
|
||||
.msg.assistant a{color:#93c5fd;text-decoration:underline}
|
||||
.msg.assistant code{
|
||||
background:rgba(0,0,0,0.35);
|
||||
padding:1px 5px;border-radius:4px;
|
||||
font-family:"SF Mono",Menlo,Monaco,monospace;
|
||||
font-size:12px;
|
||||
}
|
||||
.msg.assistant pre{
|
||||
background:rgba(0,0,0,0.40);
|
||||
padding:8px 10px;border-radius:8px;
|
||||
margin:4px 0;overflow-x:auto;
|
||||
border:1px solid rgba(255,255,255,0.08);
|
||||
}
|
||||
.msg.assistant pre code{
|
||||
background:transparent;padding:0;border-radius:0;
|
||||
font-size:11.5px;line-height:1.4;white-space:pre;
|
||||
}
|
||||
.msg.assistant strong{font-weight:600;color:#fff}
|
||||
.msg.assistant em{font-style:italic}
|
||||
.msg.assistant del{opacity:0.6;text-decoration:line-through}
|
||||
|
||||
/* Streaming caret — pulses at the end of the in-progress bubble */
|
||||
.caret{
|
||||
display:inline-block;width:6px;height:13px;
|
||||
vertical-align:text-bottom;margin-left:2px;
|
||||
background:rgba(255,255,255,0.85);
|
||||
animation:caret-blink 1s steps(1) infinite;
|
||||
}
|
||||
@keyframes caret-blink{50%{opacity:0}}
|
||||
|
||||
/* Thinking dots — shown while waiting for the first token */
|
||||
.thinking{
|
||||
display:inline-flex;gap:4px;align-items:center;padding:2px 0;
|
||||
}
|
||||
.thinking span{
|
||||
width:6px;height:6px;border-radius:50%;
|
||||
background:rgba(255,255,255,0.65);
|
||||
animation:thinking-bounce 1.2s ease-in-out infinite;
|
||||
}
|
||||
.thinking span:nth-child(2){animation-delay:0.15s}
|
||||
.thinking span:nth-child(3){animation-delay:0.30s}
|
||||
@keyframes thinking-bounce{
|
||||
0%,60%,100%{transform:translateY(0);opacity:0.4}
|
||||
30%{transform:translateY(-4px);opacity:1}
|
||||
}
|
||||
#input-bar{
|
||||
display:flex;align-items:center;gap:6px;
|
||||
border-radius:16px;padding:6px 8px;
|
||||
background:rgba(30,30,30,0.88);
|
||||
border:1px solid rgba(255,255,255,0.20);
|
||||
flex-shrink:0;
|
||||
}
|
||||
#model-wrap{
|
||||
position:relative;flex-shrink:0;
|
||||
}
|
||||
#model-select{
|
||||
background:rgba(255,255,255,0.08);color:rgba(255,255,255,0.6);
|
||||
border:none;border-radius:8px;padding:4px 22px 4px 8px;
|
||||
font-size:11px;outline:none;cursor:pointer;
|
||||
max-width:140px;
|
||||
-webkit-appearance:none;appearance:none;
|
||||
}
|
||||
#model-select:hover{background:rgba(255,255,255,0.14);color:#fff}
|
||||
#model-wrap .arrow{
|
||||
position:absolute;right:7px;top:50%;transform:translateY(-50%);
|
||||
pointer-events:none;color:rgba(255,255,255,0.35);
|
||||
}
|
||||
#model-select option,#model-select optgroup{
|
||||
background:#1e1e1e;color:#eee;
|
||||
}
|
||||
#input{
|
||||
flex:1;background:transparent;border:none;outline:none;
|
||||
color:#fff;font-size:14px;padding:6px 10px;
|
||||
}
|
||||
#input::placeholder{color:rgba(255,255,255,0.35)}
|
||||
.btn{
|
||||
display:flex;align-items:center;justify-content:center;
|
||||
width:30px;height:30px;border-radius:50%;border:none;
|
||||
background:rgba(255,255,255,0.10);
|
||||
color:rgba(255,255,255,0.6);cursor:pointer;
|
||||
transition:background .15s,color .15s;flex-shrink:0;
|
||||
}
|
||||
.btn:hover{background:rgba(255,255,255,0.20);color:#fff}
|
||||
.btn:disabled{opacity:0.25;cursor:default}
|
||||
.btn:disabled:hover{background:rgba(255,255,255,0.10)}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div id="container">
|
||||
<div id="messages"></div>
|
||||
<div id="input-bar">
|
||||
<button id="new-btn" class="btn" title="New conversation">
|
||||
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2.5" stroke-linecap="round" stroke-linejoin="round"><line x1="12" y1="5" x2="12" y2="19"/><line x1="5" y1="12" x2="19" y2="12"/></svg>
|
||||
</button>
|
||||
<div id="model-wrap">
|
||||
<select id="model-select"><option>loading...</option></select>
|
||||
<svg class="arrow" width="10" height="10" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2.5" stroke-linecap="round" stroke-linejoin="round"><polyline points="6 9 12 15 18 9"/></svg>
|
||||
</div>
|
||||
<input id="input" type="text" placeholder="Ask Jarvis anything..." autofocus>
|
||||
<button id="send-btn" class="btn" disabled title="Send">
|
||||
<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><line x1="12" y1="19" x2="12" y2="5"/><polyline points="5 12 12 5 19 12"/></svg>
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
<script type="application/json" id="saved-data">__SAVED_MESSAGES__</script>
|
||||
<script type="application/json" id="cloud-data">__CLOUD_MODELS__</script>
|
||||
<script>
|
||||
const input=document.getElementById('input');
|
||||
const sendBtn=document.getElementById('send-btn');
|
||||
const messagesEl=document.getElementById('messages');
|
||||
let streaming=false,abort=null,model='qwen3.5:4b';
|
||||
let convId='', convTitle='Overlay chat', convCreated=Date.now();
|
||||
let messages=[];
|
||||
|
||||
function genId(){return Date.now().toString(36)+Math.random().toString(36).slice(2,8)}
|
||||
|
||||
// Restore previous conversation
|
||||
try{
|
||||
const raw=document.getElementById('saved-data').textContent.trim();
|
||||
if(raw&&raw!=='__SAVED_PLACEHOLDER__'){
|
||||
const saved=JSON.parse(raw);
|
||||
if(saved.id){convId=saved.id;convTitle=saved.title||convTitle;convCreated=saved.createdAt||convCreated;messages=saved.messages||[]}
|
||||
else if(Array.isArray(saved)){messages=saved}
|
||||
}
|
||||
}catch{}
|
||||
if(!convId) convId=genId();
|
||||
if(messages.length) renderAll();
|
||||
|
||||
// Build model dropdown: local (installed) + cloud (keyed)
|
||||
const modelSelect=document.getElementById('model-select');
|
||||
let cloudModels=[];
|
||||
try{
|
||||
const cd=document.getElementById('cloud-data').textContent.trim();
|
||||
if(cd&&cd!=='__CLOUD_PLACEHOLDER__') cloudModels=JSON.parse(cd);
|
||||
}catch{}
|
||||
|
||||
fetch('/v1/models').then(r=>r.json()).then(d=>{
|
||||
const local=(Array.isArray(d)?d:(d.data||d.models||[])).map(m=>m.id||m.name).filter(Boolean);
|
||||
while(modelSelect.firstChild) modelSelect.removeChild(modelSelect.firstChild);
|
||||
if(local.length){
|
||||
const g=document.createElement('optgroup');g.label='Local';
|
||||
local.forEach(id=>{const o=document.createElement('option');o.value=id;o.textContent=id;g.appendChild(o)});
|
||||
modelSelect.appendChild(g);
|
||||
}
|
||||
if(cloudModels.length){
|
||||
const g=document.createElement('optgroup');g.label='Cloud';
|
||||
cloudModels.forEach(id=>{const o=document.createElement('option');o.value=id;o.textContent=id;g.appendChild(o)});
|
||||
modelSelect.appendChild(g);
|
||||
}
|
||||
// Restore saved model or pick first available
|
||||
const saved=messages.length&&messages[0].model;
|
||||
if(saved&&modelSelect.querySelector('option[value="'+CSS.escape(saved)+'"]')){modelSelect.value=saved}
|
||||
model=modelSelect.value||model;
|
||||
}).catch(()=>{});
|
||||
modelSelect.addEventListener('change',()=>{model=modelSelect.value});
|
||||
|
||||
const SEND='<svg width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round"><line x1="12" y1="19" x2="12" y2="5"/><polyline points="5 12 12 5 19 12"/></svg>';
|
||||
const STOP='<svg width="16" height="16" viewBox="0 0 24 24" fill="currentColor"><rect x="6" y="6" width="12" height="12" rx="2"/></svg>';
|
||||
|
||||
function renderAll(){
|
||||
while(messagesEl.firstChild) messagesEl.removeChild(messagesEl.firstChild);
|
||||
for(const m of messages) bubble(m.role,m.content);
|
||||
scroll();
|
||||
}
|
||||
function bubble(role,text){
|
||||
const d=document.createElement('div');
|
||||
d.className='msg '+(role==='user'?'user':'assistant');
|
||||
if(role==='assistant') setHtml(d, md(text||''));
|
||||
else d.textContent=text;
|
||||
messagesEl.appendChild(d);
|
||||
return d;
|
||||
}
|
||||
function scroll(){messagesEl.scrollTop=messagesEl.scrollHeight}
|
||||
|
||||
// Safe HTML injection helper. All incoming LLM/user text is escaped
|
||||
// via escHtml() before any markdown transformations, so the string
|
||||
// reaching this function only contains tags from our controlled
|
||||
// regex replacements. We use Range.createContextualFragment which is
|
||||
// the W3C-recommended way to construct a DocumentFragment from HTML.
|
||||
function setHtml(el,html){
|
||||
while(el.firstChild) el.removeChild(el.firstChild);
|
||||
const range=document.createRange();
|
||||
range.selectNodeContents(el);
|
||||
el.appendChild(range.createContextualFragment(html));
|
||||
}
|
||||
|
||||
// --- Minimal markdown renderer (inline, no deps) ---
|
||||
// Handles: fenced code, inline code, headings, bold, italic,
|
||||
// strikethrough, links, ordered/unordered lists, paragraphs.
|
||||
function escHtml(s){
|
||||
return s.replace(/&/g,'&').replace(/</g,'<').replace(/>/g,'>')
|
||||
.replace(/"/g,'"').replace(/'/g,''');
|
||||
}
|
||||
function md(src){
|
||||
if(!src)return '';
|
||||
const blocks=[],inlines=[];
|
||||
// 1. Extract fenced code blocks first so their contents aren't touched.
|
||||
src=src.replace(/```(\w*)\n?([\s\S]*?)```/g,(_,lang,code)=>{
|
||||
blocks.push({lang,code});
|
||||
return '\u0000CB'+(blocks.length-1)+'\u0000';
|
||||
});
|
||||
// 2. Extract inline code (single-line backticks).
|
||||
src=src.replace(/`([^`\n]+)`/g,(_,c)=>{
|
||||
inlines.push(c);
|
||||
return '\u0000IC'+(inlines.length-1)+'\u0000';
|
||||
});
|
||||
// 3. Escape everything else.
|
||||
src=escHtml(src);
|
||||
// 4. Headings.
|
||||
src=src.replace(/^###\s+(.+)$/gm,'<h3>$1</h3>')
|
||||
.replace(/^##\s+(.+)$/gm,'<h2>$1</h2>')
|
||||
.replace(/^#\s+(.+)$/gm,'<h1>$1</h1>');
|
||||
// 5. Bold / italic / strikethrough.
|
||||
src=src.replace(/\*\*([^*\n]+)\*\*/g,'<strong>$1</strong>')
|
||||
.replace(/__([^_\n]+)__/g,'<strong>$1</strong>')
|
||||
.replace(/(^|[^*\w])\*([^*\n]+)\*(?!\*)/g,'$1<em>$2</em>')
|
||||
.replace(/(^|[^_\w])_([^_\n]+)_(?!_)/g,'$1<em>$2</em>')
|
||||
.replace(/~~([^~\n]+)~~/g,'<del>$1</del>');
|
||||
// 6. Links — url is escaped above, so quotes are safe.
|
||||
src=src.replace(/\[([^\]\n]+)\]\(([^)\s]+)\)/g,'<a href="$2" target="_blank" rel="noopener">$1</a>');
|
||||
// 7. Unordered + ordered lists.
|
||||
src=src.replace(/(?:^[-*]\s+.+(?:\n|$))+/gm,block=>{
|
||||
const items=block.trim().split('\n')
|
||||
.map(l=>'<li>'+l.replace(/^[-*]\s+/,'')+'</li>').join('');
|
||||
return '<ul>'+items+'</ul>';
|
||||
});
|
||||
src=src.replace(/(?:^\d+\.\s+.+(?:\n|$))+/gm,block=>{
|
||||
const items=block.trim().split('\n')
|
||||
.map(l=>'<li>'+l.replace(/^\d+\.\s+/,'')+'</li>').join('');
|
||||
return '<ol>'+items+'</ol>';
|
||||
});
|
||||
// 8. Paragraphs: split on blank lines, wrap non-block chunks in <p>.
|
||||
src=src.split(/\n{2,}/).map(chunk=>{
|
||||
const t=chunk.trim();
|
||||
if(!t)return '';
|
||||
if(/^<(h\d|ul|ol|pre|blockquote)/.test(t))return t;
|
||||
if(t.startsWith('\u0000CB'))return t;
|
||||
return '<p>'+t.replace(/\n/g,'<br>')+'</p>';
|
||||
}).join('');
|
||||
// 9. Restore inline code.
|
||||
src=src.replace(/\u0000IC(\d+)\u0000/g,(_,i)=>'<code>'+escHtml(inlines[+i])+'</code>');
|
||||
// 10. Restore fenced code blocks.
|
||||
src=src.replace(/\u0000CB(\d+)\u0000/g,(_,i)=>{
|
||||
const b=blocks[+i];
|
||||
const cls=b.lang?' class="lang-'+escHtml(b.lang)+'"':'';
|
||||
return '<pre><code'+cls+'>'+escHtml(b.code)+'</code></pre>';
|
||||
});
|
||||
return src;
|
||||
}
|
||||
const THINKING='<span class="thinking"><span></span><span></span><span></span></span>';
|
||||
const CARET='<span class="caret"></span>';
|
||||
const CLOUD_PFX=['gpt-','o1-','o3-','o4-','claude-','gemini-','openrouter/','chatgpt-'];
|
||||
function save(){
|
||||
const conv={id:convId,title:convTitle,createdAt:convCreated,updatedAt:Date.now(),model,
|
||||
messages:messages.map((m,i)=>{
|
||||
const o={id:convId+'_'+i,role:m.role,content:m.content,timestamp:m.timestamp||Date.now()};
|
||||
if(m.usage)o.usage=m.usage;
|
||||
if(m.telemetry)o.telemetry=m.telemetry;
|
||||
return o;
|
||||
})};
|
||||
try{window.webkit.messageHandlers.overlay.postMessage('save:'+JSON.stringify(conv))}catch{}
|
||||
}
|
||||
|
||||
input.addEventListener('input',()=>{sendBtn.disabled=!input.value.trim()||streaming});
|
||||
input.addEventListener('keydown',e=>{
|
||||
if(e.key==='Enter'&&!e.shiftKey){e.preventDefault();send()}
|
||||
});
|
||||
document.addEventListener('keydown',e=>{
|
||||
if(e.key==='Escape'){
|
||||
if(streaming){abort&&abort.abort()}
|
||||
else{try{window.webkit.messageHandlers.overlay.postMessage('hide')}catch{}}
|
||||
}
|
||||
});
|
||||
sendBtn.addEventListener('click',()=>{if(streaming){abort&&abort.abort()}else send()});
|
||||
document.getElementById('new-btn').addEventListener('click',()=>{
|
||||
convId=genId();convTitle='Overlay chat';convCreated=Date.now();
|
||||
messages=[];
|
||||
while(messagesEl.firstChild) messagesEl.removeChild(messagesEl.firstChild);
|
||||
save();input.focus();
|
||||
});
|
||||
|
||||
async function send(){
|
||||
const text=input.value.trim();
|
||||
if(!text||streaming)return;
|
||||
input.value='';sendBtn.disabled=true;
|
||||
|
||||
messages.push({role:'user',content:text,timestamp:Date.now()});
|
||||
if(messages.length===1) convTitle=text.slice(0,50)+(text.length>50?'...':'');
|
||||
bubble('user',text);scroll();save();
|
||||
|
||||
streaming=true;abort=new AbortController();
|
||||
setHtml(sendBtn,STOP);sendBtn.disabled=false;
|
||||
const b=bubble('assistant','');
|
||||
// Show thinking dots until the first token arrives.
|
||||
setHtml(b,THINKING);
|
||||
scroll();
|
||||
|
||||
let acc='',usage=null,complexity=null,ttft=0;
|
||||
const t0=Date.now();
|
||||
// Throttle markdown re-renders to ~30fps so tight streams don't
|
||||
// rebuild the DOM on every single token.
|
||||
let pending=false;
|
||||
const render=()=>{
|
||||
if(pending)return;
|
||||
pending=true;
|
||||
requestAnimationFrame(()=>{
|
||||
pending=false;
|
||||
setHtml(b, md(acc)+CARET);
|
||||
scroll();
|
||||
});
|
||||
};
|
||||
try{
|
||||
const r=await fetch('/v1/chat/completions',{
|
||||
method:'POST',headers:{'Content-Type':'application/json'},
|
||||
body:JSON.stringify({model,messages,stream:true}),
|
||||
signal:abort.signal
|
||||
});
|
||||
if(!r.ok)throw new Error(r.status);
|
||||
const reader=r.body.getReader(),dec=new TextDecoder();
|
||||
let buf='';
|
||||
for(;;){
|
||||
const{done,value}=await reader.read();
|
||||
if(done)break;
|
||||
buf+=dec.decode(value,{stream:true});
|
||||
const lines=buf.split('\n');buf=lines.pop()||'';
|
||||
for(const ln of lines){
|
||||
if(!ln.startsWith('data: '))continue;
|
||||
const d=ln.slice(6);if(d==='[DONE]')break;
|
||||
try{
|
||||
const p=JSON.parse(d);
|
||||
if(p.usage)usage=p.usage;
|
||||
if(p.complexity)complexity=p.complexity;
|
||||
const c=p.choices?.[0]?.delta?.content;
|
||||
if(c){if(!ttft)ttft=Date.now()-t0;acc+=c;render()}
|
||||
}catch{}
|
||||
}
|
||||
}
|
||||
}catch(e){
|
||||
if(e.name!=='AbortError'){
|
||||
acc='Could not get a response. Is the backend running?';
|
||||
b.textContent=acc;
|
||||
}
|
||||
}finally{
|
||||
streaming=false;abort=null;
|
||||
setHtml(sendBtn,SEND);sendBtn.disabled=!input.value.trim();
|
||||
input.focus();
|
||||
// Final render without the caret.
|
||||
if(acc) setHtml(b, md(acc));
|
||||
else if(b.querySelector('.thinking')) b.textContent='';
|
||||
}
|
||||
if(acc){
|
||||
const totalMs=Date.now()-t0;
|
||||
const engine=CLOUD_PFX.some(p=>model.startsWith(p))?'cloud':'ollama';
|
||||
const telem={engine,model_id:model,total_ms:totalMs,ttft_ms:ttft||undefined,
|
||||
tokens_per_sec:usage?.completion_tokens?usage.completion_tokens/(totalMs/1000):undefined,
|
||||
complexity_score:complexity?.score,complexity_tier:complexity?.tier,
|
||||
suggested_max_tokens:complexity?.suggested_max_tokens};
|
||||
messages.push({role:'assistant',content:acc,timestamp:Date.now(),usage:usage||undefined,telemetry:telem});
|
||||
save();
|
||||
}
|
||||
}
|
||||
|
||||
window.addEventListener('focus',()=>input.focus());
|
||||
|
||||
// --- Drag-to-move: click anywhere to drag, with threshold so
|
||||
// clicks on inputs/buttons still work normally ---
|
||||
(function(){
|
||||
let down=false,dragging=false,sx=0,sy=0;
|
||||
const THRESH=3;
|
||||
const INTERACTIVE='input,select,button,textarea,a,label,[contenteditable="true"],[role="button"]';
|
||||
document.addEventListener('mousedown',e=>{
|
||||
if(e.target.closest('select'))return;
|
||||
down=true;dragging=false;sx=e.screenX;sy=e.screenY;
|
||||
});
|
||||
// macOS native select menus swallow the mouseup — reset on change too
|
||||
document.getElementById('model-select').addEventListener('mousedown',()=>{down=false;dragging=false});
|
||||
document.getElementById('model-select').addEventListener('change',()=>{down=false;dragging=false});
|
||||
document.addEventListener('mousemove',e=>{
|
||||
if(!down)return;
|
||||
// If the mouse button is no longer pressed (e.g. the user released
|
||||
// it over a native menu that swallowed mouseup), abort the drag.
|
||||
if(e.buttons===0){down=false;dragging=false;return;}
|
||||
const dx=e.screenX-sx,dy=e.screenY-sy;
|
||||
if(!dragging){
|
||||
if(Math.abs(dx)+Math.abs(dy)<THRESH)return;
|
||||
dragging=true;
|
||||
}
|
||||
sx=e.screenX;sy=e.screenY;
|
||||
try{window.webkit.messageHandlers.overlay.postMessage('drag:'+dx+','+dy)}catch{}
|
||||
});
|
||||
const reset=()=>{down=false;dragging=false};
|
||||
document.addEventListener('mouseup',reset);
|
||||
window.addEventListener('blur',reset);
|
||||
document.addEventListener('mouseleave',reset);
|
||||
})();
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -553,3 +553,60 @@ metric can be improved without degrading another.
|
||||
The optimization framework has full Rust parity via the `openjarvis-learning`
|
||||
crate, with PyO3 bindings exposing `OptimizationStore` and `LLMOptimizer`
|
||||
to Python.
|
||||
|
||||
---
|
||||
|
||||
## Distillation (Frontier-Driven Harness Learning)
|
||||
|
||||
The distillation subsystem uses a frontier closed-source model (the "teacher") as a meta-engineer for the local student's full harness — not just its weights. Instead of pushing knowledge into a small model's weights, we push a frontier model's engineering judgement into the surrounding configuration: prompts, routing, agent class, tool availability, and tool descriptions.
|
||||
|
||||
### Where it lives
|
||||
|
||||
`learning/distillation/` is the fifth subsystem within the Learning pillar, alongside `learning/routing/`, `learning/optimize/`, `learning/training/`, and `learning/intelligence/`.
|
||||
|
||||
### Four-phase loop
|
||||
|
||||
```
|
||||
Trigger → Diagnose → Plan → Execute → Record
|
||||
```
|
||||
|
||||
1. **Diagnose** — TeacherAgent (frontier model with diagnostic tools) analyzes traces, runs student/teacher comparisons, identifies 2-5 failure clusters with evidence.
|
||||
2. **Plan** — LearningPlanner converts diagnosis into a typed LearningPlan with deterministic risk tier assignment and patch/replace downgrade.
|
||||
3. **Execute** — Per-edit loop: EditApplier validates + applies, BenchmarkGate scores, CheckpointStore commits or rolls back.
|
||||
4. **Record** — LearningSession persisted to SQLite + JSON artifact.
|
||||
|
||||
### Key components
|
||||
|
||||
| Component | Module | Purpose |
|
||||
|-----------|--------|---------|
|
||||
| `DistillationOrchestrator` | `orchestrator.py` | Top-level session driver |
|
||||
| `TeacherAgent` | `diagnose/teacher_agent.py` | Frontier model tool-calling loop |
|
||||
| `DiagnosisRunner` | `diagnose/runner.py` | Phase 1 orchestration |
|
||||
| `LearningPlanner` | `plan/planner.py` | Diagnosis → typed LearningPlan |
|
||||
| `EditApplier` + registry | `execute/base.py` | Abstract applier interface |
|
||||
| `BenchmarkGate` | `gate/benchmark_gate.py` | Benchmark-based accept/reject |
|
||||
| `CheckpointStore` | `checkpoint/store.py` | Git-backed config rollback |
|
||||
| `SessionStore` | `storage/session_store.py` | SQLite session persistence |
|
||||
|
||||
### Relationship to existing subsystems
|
||||
|
||||
| Existing | Relationship |
|
||||
|----------|-------------|
|
||||
| `LearningOrchestrator` | Sibling — stays untouched |
|
||||
| `LLMOptimizer` / `OptimizationStore` | Sibling — grid-search vs root-cause |
|
||||
| `LearnedRouterPolicy` | Reused — routing edits update it |
|
||||
| `TraceJudge` | Reused — benchmark scoring |
|
||||
| `PersonalBenchmarkSynthesizer` | Extended — auto-refresh, gold answers |
|
||||
| `TraceStore` | Reused — read-only access from diagnostic tools |
|
||||
|
||||
### Risk tier system
|
||||
|
||||
Every edit is assigned a tier from a deterministic lookup table:
|
||||
|
||||
| Tier | Ops | Behavior |
|
||||
|------|-----|----------|
|
||||
| `auto` | Model routing/params, tool add/remove/description, agent params | Apply if gate passes |
|
||||
| `review` | System prompt edits, agent class, few-shot exemplars | Queue for user approval |
|
||||
| `manual` | LoRA fine-tuning (v2) | Never auto-apply |
|
||||
|
||||
See [Distillation user guide](../user-guide/learning-distillation.md) for CLI usage and configuration.
|
||||
|
||||
@@ -45,7 +45,7 @@ The memory pipeline includes document ingestion, chunking, embedding generation,
|
||||
|
||||
The Learning system is the fifth primitive, connecting the other four through **trace-driven feedback**. Every agent interaction can produce a `Trace` capturing the full sequence of steps — routing decisions, memory retrieval, inference calls, tool invocations, and final responses. The `TraceAnalyzer` computes statistics from accumulated traces, and the `TraceDrivenPolicy` uses these statistics to learn which model/agent/tool combinations produce the best outcomes for different query types.
|
||||
|
||||
The learning system is configured through nested sub-sections in `config.toml`: `[learning.routing]` controls the router policy (heuristic, learned, sft, grpo), `[learning.intelligence]` controls the model-level learning policy, `[learning.agent]` controls agent advisor and ICL updater policies, and `[learning.metrics]` sets the composite reward function weights.
|
||||
The learning system is configured through nested sub-sections in `config.toml`: `[learning.routing]` controls the router policy (heuristic, learned, sft, grpo), `[learning.intelligence]` controls the model-level learning policy, `[learning.agent]` controls agent advisor and ICL updater policies, and `[learning.metrics]` sets the composite reward function weights. The pillar also includes the distillation subsystem, a frontier-driven loop that improves the local harness — see [Learning architecture: Distillation](learning.md#distillation-frontier-driven-harness-learning).
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ These are the areas where active development is happening and contributions are
|
||||
- **Energy-aware routing** — using power consumption data from telemetry to optimize for energy efficiency alongside latency and quality
|
||||
- **Plugin ecosystem** — community-contributed engines, tools, and agents distributed as Python packages
|
||||
- **Federated memory** — memory backends that synchronize across devices
|
||||
- **Distillation:** Frontier-driven harness learning — a frontier model analyzes your traces and proposes config improvements. See [user guide](../user-guide/learning-distillation.md) and [architecture](../architecture/learning.md#distillation-frontier-driven-harness-learning).
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -1057,3 +1057,59 @@ OpenJarvis respects the following environment variables:
|
||||
- [Architecture Overview](../architecture/overview.md) — Understand how the pieces fit together
|
||||
- [Intelligence Primitive](../architecture/intelligence.md) — Model identity and generation defaults
|
||||
- [Learning & Traces](../architecture/learning.md) — Router policies and the trace-driven feedback loop
|
||||
|
||||
---
|
||||
|
||||
## Learning & Distillation
|
||||
|
||||
The distillation subsystem uses a frontier model to automatically improve your local agent configuration. See the [user guide](../user-guide/learning-distillation.md) for a full walkthrough.
|
||||
|
||||
### `[learning.distillation]`
|
||||
|
||||
| Key | Type | Default | Description |
|
||||
|-----|------|---------|-------------|
|
||||
| `enabled` | bool | `true` | Gate the entire distillation subsystem |
|
||||
| `autonomy_mode` | string | `"tiered"` | `auto`, `tiered`, or `manual` |
|
||||
| `teacher_model` | string | `"claude-opus-4-6"` | Frontier model for diagnosis and planning |
|
||||
| `max_cost_per_session_usd` | float | `5.0` | Per-session teacher API budget |
|
||||
| `max_tool_calls_per_diagnosis` | int | `30` | Max teacher tool calls in diagnosis phase |
|
||||
|
||||
### `[learning.distillation.triggers]`
|
||||
|
||||
| Key | Type | Default | Description |
|
||||
|-----|------|---------|-------------|
|
||||
| `scheduled_enabled` | bool | `true` | Enable daily scheduled sessions |
|
||||
| `scheduled_cron` | string | `"0 3 * * *"` | Cron expression for scheduled trigger |
|
||||
| `scheduled_min_new_traces` | int | `20` | Minimum new traces to trigger |
|
||||
| `cluster_enabled` | bool | `true` | Enable failure cluster trigger |
|
||||
| `cluster_check_interval_minutes` | int | `60` | How often to check for clusters |
|
||||
| `cluster_min_size` | int | `5` | Minimum traces in a cluster |
|
||||
| `cluster_failure_threshold` | float | `0.3` | Feedback <= this counts as failure |
|
||||
|
||||
### `[learning.distillation.gate]`
|
||||
|
||||
| Key | Type | Default | Description |
|
||||
|-----|------|---------|-------------|
|
||||
| `min_improvement` | float | `0.0` | Minimum overall score improvement to accept |
|
||||
| `max_regression` | float | `0.05` | Maximum per-cluster score drop before rejecting |
|
||||
| `benchmark_subsample_size` | int | `50` | Tasks per gate run |
|
||||
| `full_benchmark` | bool | `false` | Disable subsampling (slower, more accurate) |
|
||||
|
||||
### `[learning.distillation.benchmark]`
|
||||
|
||||
| Key | Type | Default | Description |
|
||||
|-----|------|---------|-------------|
|
||||
| `synthesis_feedback_threshold` | float | `0.7` | Min feedback for benchmark traces |
|
||||
| `max_benchmark_size` | int | `200` | Max tasks in the benchmark |
|
||||
| `auto_refresh` | bool | `true` | Auto-mine new high-feedback traces |
|
||||
| `max_synthesis_cost_usd_per_refresh` | float | `2.0` | Cost cap per benchmark refresh |
|
||||
|
||||
### `[learning.distillation.tier_overrides]`
|
||||
|
||||
Override the default risk tier for any operation. Keys are operation names, values are tier strings (`auto`, `review`, `manual`).
|
||||
|
||||
```toml
|
||||
[learning.distillation.tier_overrides]
|
||||
# patch_system_prompt = "auto" # promote to auto after trust
|
||||
# replace_system_prompt = "auto"
|
||||
```
|
||||
|
||||
@@ -431,3 +431,101 @@ curl http://localhost:8000/v1/chat/completions \
|
||||
```
|
||||
|
||||
When an agent is configured (e.g., `--agent orchestrator`), non-streaming requests are routed through the agent with access to all registered tools. For tool-capable agents (`orchestrator`, `react`, `openhands`), all registered tools are automatically loaded and made available.
|
||||
|
||||
---
|
||||
|
||||
## `jarvis learning`
|
||||
|
||||
Frontier-driven harness learning (distillation). Manages learning sessions, reviews pending edits, and controls the benchmark gate.
|
||||
|
||||
### `jarvis learning init`
|
||||
|
||||
Initialize the distillation checkpoint repo and directory layout.
|
||||
|
||||
```bash
|
||||
jarvis learning init
|
||||
```
|
||||
|
||||
### `jarvis learning run`
|
||||
|
||||
Run an on-demand learning session.
|
||||
|
||||
```bash
|
||||
jarvis learning run
|
||||
jarvis learning run --autonomy auto # auto-apply all edits
|
||||
jarvis learning run --autonomy manual # dry-run, everything goes to review
|
||||
```
|
||||
|
||||
| Flag | Default | Description |
|
||||
|------|---------|-------------|
|
||||
| `--autonomy` | `tiered` | `auto`, `tiered`, or `manual` |
|
||||
|
||||
### `jarvis learning history`
|
||||
|
||||
List past learning sessions.
|
||||
|
||||
```bash
|
||||
jarvis learning history
|
||||
jarvis learning history --limit 5
|
||||
```
|
||||
|
||||
### `jarvis learning show`
|
||||
|
||||
Show details of a learning session (diagnosis, plan, outcomes, cost).
|
||||
|
||||
```bash
|
||||
jarvis learning show <session-id>
|
||||
```
|
||||
|
||||
### `jarvis learning review`
|
||||
|
||||
List all pending edits awaiting approval.
|
||||
|
||||
```bash
|
||||
jarvis learning review
|
||||
```
|
||||
|
||||
### `jarvis learning approve`
|
||||
|
||||
Approve a pending edit (still goes through the benchmark gate).
|
||||
|
||||
```bash
|
||||
jarvis learning approve <edit-id>
|
||||
```
|
||||
|
||||
### `jarvis learning reject`
|
||||
|
||||
Reject a pending edit.
|
||||
|
||||
```bash
|
||||
jarvis learning reject <edit-id>
|
||||
jarvis learning reject <edit-id> --reason "too aggressive"
|
||||
```
|
||||
|
||||
### `jarvis learning rollback`
|
||||
|
||||
Rollback a session's committed edits (creates revert commits).
|
||||
|
||||
```bash
|
||||
jarvis learning rollback <session-id>
|
||||
jarvis learning rollback --last
|
||||
```
|
||||
|
||||
### `jarvis learning benchmark`
|
||||
|
||||
Personal benchmark management.
|
||||
|
||||
```bash
|
||||
jarvis learning benchmark show # current stats
|
||||
jarvis learning benchmark refresh # manual refresh
|
||||
```
|
||||
|
||||
### `jarvis learning daemon`
|
||||
|
||||
Background learning daemon.
|
||||
|
||||
```bash
|
||||
jarvis learning daemon start
|
||||
jarvis learning daemon stop
|
||||
jarvis learning daemon status
|
||||
```
|
||||
|
||||
@@ -14,6 +14,8 @@ The OpenJarvis evaluation framework (`openjarvis-evals`) measures model **correc
|
||||
|
||||
---
|
||||
|
||||
> **Tip:** The distillation system uses this same eval infrastructure to gate edits against your personal benchmark. See [Learning & Distillation](learning-distillation.md).
|
||||
|
||||
## Installation
|
||||
|
||||
The evaluation framework is a standalone package in the `evals/` directory. Install it alongside OpenJarvis:
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
# Learning & Distillation
|
||||
|
||||
Use a frontier model as a meta-engineer to automatically improve your local agent's prompts, routing, and tools — reversibly, with benchmark-gated quality control.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Initialize
|
||||
|
||||
```bash
|
||||
jarvis learning init
|
||||
```
|
||||
|
||||
This creates the distillation directory layout under `~/.openjarvis/learning/` and initializes a git checkpoint repo at `~/.openjarvis/.git` for tracking config changes.
|
||||
|
||||
### 2. Run your first session
|
||||
|
||||
Once you have at least 20 traces from regular use:
|
||||
|
||||
```bash
|
||||
jarvis learning run
|
||||
```
|
||||
|
||||
The system will:
|
||||
1. **Diagnose** — analyze your traces using a frontier model
|
||||
2. **Plan** — propose typed edits to your config
|
||||
3. **Execute** — apply edits that pass the benchmark gate
|
||||
4. **Record** — persist the session for history and rollback
|
||||
|
||||
### 3. Check results
|
||||
|
||||
```bash
|
||||
jarvis learning history
|
||||
jarvis learning show <session-id>
|
||||
```
|
||||
|
||||
## How a Learning Session Works
|
||||
|
||||
A learning session has four phases:
|
||||
|
||||
### Phase 1: Diagnose
|
||||
|
||||
A frontier model (the "teacher", default `claude-opus-4-6`) analyzes your recent traces using read-only diagnostic tools. It identifies **failure clusters** — groups of related failures with shared root causes. The teacher must actually re-run your student on sample tasks and compare outputs to populate failure rates. This forces evidence-based diagnosis.
|
||||
|
||||
**Output:** `diagnosis.md` with narrative analysis + structured failure clusters.
|
||||
|
||||
### Phase 2: Plan
|
||||
|
||||
A second teacher call converts the diagnosis into a typed `LearningPlan` — a list of `Edit` objects, each targeting a specific part of your configuration (model routing, system prompts, tool availability, etc.). The teacher cannot pick risk tiers — those are assigned deterministically from a lookup table.
|
||||
|
||||
**Output:** `plan.json` frozen and immutable.
|
||||
|
||||
### Phase 3: Execute
|
||||
|
||||
Each edit is applied through its registered `EditApplier`, then scored against your personal benchmark. Edits that improve the benchmark are committed; edits that cause regressions are rolled back. Edits in the `review` tier are queued for your approval instead of being auto-applied.
|
||||
|
||||
**Output:** Git commits in the checkpoint repo + `EditOutcome` records.
|
||||
|
||||
### Phase 4: Record
|
||||
|
||||
The session is persisted to `learning.db` (SQLite index) and `session.json` (authoritative artifact). You can query history, show details, and rollback any session.
|
||||
|
||||
## Configuration
|
||||
|
||||
Add to `~/.openjarvis/config.toml`:
|
||||
|
||||
```toml
|
||||
[learning.distillation]
|
||||
enabled = true # gate the entire subsystem
|
||||
autonomy_mode = "tiered" # auto | tiered | manual
|
||||
teacher_model = "claude-opus-4-6" # any CloudEngine-supported model
|
||||
max_cost_per_session_usd = 5.0 # per-session teacher API budget
|
||||
max_tool_calls_per_diagnosis = 30 # max teacher tool calls in diagnosis
|
||||
```
|
||||
|
||||
### Trigger configuration
|
||||
|
||||
```toml
|
||||
[learning.distillation.triggers]
|
||||
scheduled_enabled = true
|
||||
scheduled_cron = "0 3 * * *" # daily at 03:00 local
|
||||
scheduled_min_new_traces = 20 # minimum new traces to trigger
|
||||
|
||||
cluster_enabled = true
|
||||
cluster_check_interval_minutes = 60
|
||||
cluster_min_size = 5
|
||||
cluster_failure_threshold = 0.3 # feedback <= this counts as failure
|
||||
```
|
||||
|
||||
### Gate configuration
|
||||
|
||||
```toml
|
||||
[learning.distillation.gate]
|
||||
min_improvement = 0.0 # any improvement accepted (raise for margin)
|
||||
max_regression = 0.05 # max per-cluster score drop
|
||||
benchmark_subsample_size = 50 # tasks per gate run
|
||||
full_benchmark = false # set true to disable subsampling
|
||||
```
|
||||
|
||||
### Benchmark configuration
|
||||
|
||||
```toml
|
||||
[learning.distillation.benchmark]
|
||||
synthesis_feedback_threshold = 0.7 # min feedback for benchmark traces
|
||||
max_benchmark_size = 200 # max tasks in the benchmark
|
||||
auto_refresh = true # auto-mine new high-feedback traces
|
||||
max_synthesis_cost_usd_per_refresh = 2.0 # separate from session budget
|
||||
```
|
||||
|
||||
### Risk tier overrides
|
||||
|
||||
Power users can override the default tier for any operation:
|
||||
|
||||
```toml
|
||||
[learning.distillation.tier_overrides]
|
||||
# Promote prompt edits to auto-apply after trust is established:
|
||||
# patch_system_prompt = "auto"
|
||||
# replace_system_prompt = "auto"
|
||||
```
|
||||
|
||||
## Risk Tiers
|
||||
|
||||
Every edit is assigned a risk tier that controls how it's applied:
|
||||
|
||||
| Tier | Behavior | Default ops |
|
||||
|------|----------|-------------|
|
||||
| **auto** | Applied automatically if benchmark gate passes | Model routing, model params, tool add/remove/description, agent params |
|
||||
| **review** | Queued for user approval in `jarvis learning review` | System prompt edits, agent class changes, few-shot exemplars |
|
||||
| **manual** | Never auto-applied; requires explicit approval | LoRA fine-tuning (v2) |
|
||||
|
||||
The tier is assigned deterministically from the edit operation — the teacher cannot override it.
|
||||
|
||||
## Reviewing Edits
|
||||
|
||||
When edits land in the review queue:
|
||||
|
||||
```bash
|
||||
# List all pending edits
|
||||
jarvis learning review
|
||||
|
||||
# Approve an edit (still goes through the benchmark gate)
|
||||
jarvis learning approve <edit-id>
|
||||
|
||||
# Reject an edit with a reason
|
||||
jarvis learning reject <edit-id> --reason "prompt change too aggressive"
|
||||
```
|
||||
|
||||
Even approved edits are gated by the benchmark — approval means "try it", not "force it".
|
||||
|
||||
## Rollback and History
|
||||
|
||||
Every edit creates a git commit in the checkpoint repo at `~/.openjarvis/.git`. This is separate from your OpenJarvis source repo.
|
||||
|
||||
```bash
|
||||
# List past sessions
|
||||
jarvis learning history --limit 20
|
||||
|
||||
# Show session details (diagnosis, plan, outcomes, cost)
|
||||
jarvis learning show <session-id>
|
||||
|
||||
# Rollback a session (creates new revert commits, preserves history)
|
||||
jarvis learning rollback <session-id>
|
||||
jarvis learning rollback --last
|
||||
```
|
||||
|
||||
Rollback never rewrites git history — it creates new revert commits so the audit trail stays intact.
|
||||
|
||||
## Cost Controls
|
||||
|
||||
Three cost boundaries prevent runaway spending:
|
||||
|
||||
1. **`max_cost_per_session_usd`** (default $5.00) — caps the total teacher API cost per session (diagnosis + planning).
|
||||
2. **`max_synthesis_cost_usd_per_refresh`** (default $2.00) — caps the cost of generating gold answers for new benchmark tasks. Separate from the session budget.
|
||||
3. **`teacher_model`** — choose a cheaper model (e.g., `claude-sonnet-4-6`) to reduce per-token costs at the expense of diagnosis quality.
|
||||
|
||||
Cost is tracked on every `LearningSession` as `teacher_cost_usd` and surfaced in `jarvis learning show`.
|
||||
|
||||
## The Personal Benchmark
|
||||
|
||||
The benchmark is your acceptance gate's source of truth — a set of tasks distilled from your high-quality traces, scored by an LLM-as-judge against frontier gold answers.
|
||||
|
||||
**How it's built:**
|
||||
1. Traces with feedback >= 0.7 are candidates
|
||||
2. Tasks are grouped by query class and deduplicated
|
||||
3. For each task, the teacher generates a gold reference answer
|
||||
4. The benchmark is versioned (`personal_v1.json`, `personal_v2.json`, ...)
|
||||
|
||||
**Auto-refresh:** The benchmark grows over time as you accumulate more traces. New tasks are added automatically during background refresh cycles.
|
||||
|
||||
```bash
|
||||
# Manual refresh
|
||||
jarvis learning benchmark refresh
|
||||
|
||||
# Show stats
|
||||
jarvis learning benchmark show
|
||||
```
|
||||
|
||||
## Cold Start: What to Expect on Day One
|
||||
|
||||
The system needs real usage data before it can learn:
|
||||
|
||||
- **< 20 traces:** `jarvis learning run` returns "Not enough traces yet." Triggers are no-ops.
|
||||
- **20+ traces, < 10 high-feedback:** Enough for diagnosis, but no benchmark yet. Sessions will run diagnosis but can't gate edits.
|
||||
- **10+ high-feedback traces:** Bootstrap benchmark is created automatically (`personal_v1.json`). Full learning loop is available.
|
||||
|
||||
**Getting there faster:** Use OpenJarvis normally and provide feedback on results (thumbs up/down in the UI, or `jarvis feedback` in the CLI).
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
| Error | Cause | Fix |
|
||||
|-------|-------|-----|
|
||||
| "Not enough traces yet" | Fewer than 20 traces in the store | Use OpenJarvis more, provide feedback |
|
||||
| "Working tree dirty, cannot stage" | Manual edits to `~/.openjarvis/config.toml` during a session | Commit or revert manual changes first |
|
||||
| "All clusters dropped: insufficient evidence" | Teacher diagnosed clusters but couldn't reproduce failures | Check that the student is actually failing on the flagged tasks |
|
||||
| "ConfigurationError: distillation root inside source tree" | `OPENJARVIS_HOME` points inside the repo | Set `OPENJARVIS_HOME` to `~/.openjarvis` (default) or another external dir |
|
||||
| "Personal benchmark is empty" | Not enough high-feedback traces yet | Provide feedback on 10+ traces with score >= 0.7 |
|
||||
|
||||
## Where Artifacts Live
|
||||
|
||||
All distillation artifacts live under `~/.openjarvis/` (never inside the source repo):
|
||||
|
||||
```
|
||||
~/.openjarvis/
|
||||
├── config.toml # Your configuration (git-tracked by checkpoint)
|
||||
├── agents/ # Agent prompts (git-tracked)
|
||||
├── tools/ # Tool descriptions (git-tracked)
|
||||
├── .git/ # Checkpoint repo for rollback
|
||||
└── learning/
|
||||
├── learning.db # SQLite session index
|
||||
├── benchmarks/ # Personal benchmark versions + gold answers
|
||||
├── sessions/ # Per-session artifacts (diagnosis, plan, traces)
|
||||
└── pending_review/ # Edits awaiting user approval
|
||||
```
|
||||
|
||||
## Background Daemon
|
||||
|
||||
For continuous learning:
|
||||
|
||||
```bash
|
||||
jarvis learning daemon start # Start background watcher
|
||||
jarvis learning daemon status # Check if running
|
||||
jarvis learning daemon stop # Stop the daemon
|
||||
```
|
||||
|
||||
The daemon runs the scheduled trigger (default: daily at 03:00) and the cluster trigger (watches for failure patterns in real-time).
|
||||
|
||||
## See Also
|
||||
|
||||
- [Architecture: Learning](../architecture/learning.md#distillation-frontier-driven-harness-learning) — internal architecture of the distillation subsystem
|
||||
- [User Guide: Evaluations](evaluations.md) — the eval infrastructure that powers the benchmark gate
|
||||
- [User Guide: CLI](cli.md#jarvis-learning) — full CLI reference
|
||||
- [Getting Started: Configuration](../getting-started/configuration.md) — all config knobs
|
||||
@@ -180,5 +180,6 @@ nav:
|
||||
- Scheduler: user-guide/scheduler.md
|
||||
- Telemetry: user-guide/telemetry.md
|
||||
- Security: user-guide/security.md
|
||||
- Learning & Distillation: user-guide/learning-distillation.md
|
||||
- Leaderboard: leaderboard.md
|
||||
- Roadmap: development/roadmap.md
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
#!/usr/bin/env python3
|
||||
"""A1: Seed feedback on all traces in the local M1 traces.db.
|
||||
|
||||
- Judges each unscored trace with Sonnet 4.6 using the calibration-validated prompt
|
||||
- Parallelized via ThreadPoolExecutor (8 workers) for I/O-bound API calls
|
||||
- Writes feedback to local traces.db via TraceStore.update_feedback
|
||||
- Logs every call to a JSONL audit file
|
||||
- Idempotent: skips traces that already have feedback (safe to re-run)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import sqlite3
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.core.types import Message, Role
|
||||
from openjarvis.engine.cloud import CloudEngine
|
||||
from openjarvis.traces.store import TraceStore
|
||||
|
||||
HOME = Path(os.environ.get("OPENJARVIS_HOME", "/scratch/user/jonsaadfalcon/openjarvis-m1"))
|
||||
DB = HOME / "traces.db"
|
||||
LOG = HOME / "a1_feedback_log.jsonl"
|
||||
MODEL = "claude-sonnet-4-6"
|
||||
MAX_WORKERS = 8
|
||||
|
||||
JUDGE_PROMPT = """\
|
||||
You are evaluating whether an AI agent successfully completed its assigned task.
|
||||
|
||||
Assign a SCORE from the set {{0.2, 0.4, 0.6, 0.8}} using this rubric:
|
||||
|
||||
- 0.8 = Clean success. Task completed correctly. Minor stylistic issues don't affect correctness.
|
||||
- 0.6 = Partial. Real progress made but the answer has real gaps — missed requirement, incomplete output, recovered from errors but final result imperfect.
|
||||
- 0.4 = Poor. Some progress but the result is clearly incomplete, wrong, or the agent got mostly stuck.
|
||||
- 0.2 = Failure. Agent crashed, got stuck in a loop, gave up, hit a budget/poll/token limit before finishing, or produced no usable result.
|
||||
|
||||
IMPORTANT: Do not trust the agent's own self-assessment. Agents often narrate "I am stuck" or "I hit an error" — those are failure signals. Agents sometimes claim success when the actual output is incomplete — look at the concrete result, not the rhetoric.
|
||||
|
||||
TASK QUERY (first 1200 chars):
|
||||
<<<
|
||||
{query}
|
||||
>>>
|
||||
|
||||
AGENT FINAL RESULT (first 2500 chars):
|
||||
<<<
|
||||
{result_head}
|
||||
>>>
|
||||
{tail_section}
|
||||
Respond in EXACTLY this format, nothing else:
|
||||
SCORE=<one of 0.2, 0.4, 0.6, 0.8>
|
||||
REASON=<one brief sentence>
|
||||
"""
|
||||
|
||||
|
||||
SCORE_RE = re.compile(r"SCORE\s*=\s*(0?\.[2468])", re.IGNORECASE)
|
||||
REASON_RE = re.compile(r"REASON\s*=\s*(.+?)\s*$", re.IGNORECASE | re.DOTALL)
|
||||
|
||||
|
||||
def build_prompt(query: str, result: str) -> str:
|
||||
q = (query or "")[:1200]
|
||||
head = (result or "")[:2500]
|
||||
if result and len(result) > 3000:
|
||||
tail = f"\nAGENT FINAL RESULT (last 500 chars):\n<<<\n{result[-500:]}\n>>>\n"
|
||||
else:
|
||||
tail = ""
|
||||
return JUDGE_PROMPT.format(query=q, result_head=head, tail_section=tail)
|
||||
|
||||
|
||||
def judge_one(ce: CloudEngine, trace_id: str, query: str, result: str) -> dict:
|
||||
prompt = build_prompt(query, result)
|
||||
t0 = time.time()
|
||||
try:
|
||||
resp = ce.generate(
|
||||
messages=[Message(role=Role.USER, content=prompt)],
|
||||
model=MODEL,
|
||||
max_tokens=150,
|
||||
temperature=0.0,
|
||||
)
|
||||
content = resp.get("content", "") or ""
|
||||
cost = resp.get("cost_usd", 0.0) or 0.0
|
||||
usage = resp.get("usage", {}) or {}
|
||||
in_tok = usage.get("prompt_tokens", 0) or usage.get("input_tokens", 0)
|
||||
out_tok = usage.get("completion_tokens", 0) or usage.get("output_tokens", 0)
|
||||
|
||||
m = SCORE_RE.search(content)
|
||||
score = float(m.group(1)) if m else None
|
||||
mr = REASON_RE.search(content)
|
||||
reason = mr.group(1).strip() if mr else "(parse failed)"
|
||||
|
||||
return {
|
||||
"trace_id": trace_id,
|
||||
"score": score,
|
||||
"reason": reason,
|
||||
"raw": content,
|
||||
"cost": cost,
|
||||
"input_tokens": in_tok,
|
||||
"output_tokens": out_tok,
|
||||
"elapsed": time.time() - t0,
|
||||
"judged_at": datetime.utcnow().isoformat(timespec="seconds") + "Z",
|
||||
"error": None,
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"trace_id": trace_id,
|
||||
"score": None,
|
||||
"reason": None,
|
||||
"raw": None,
|
||||
"cost": 0.0,
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
"elapsed": time.time() - t0,
|
||||
"judged_at": datetime.utcnow().isoformat(timespec="seconds") + "Z",
|
||||
"error": f"{type(e).__name__}: {e}",
|
||||
}
|
||||
|
||||
|
||||
def main() -> int:
|
||||
if not os.environ.get("ANTHROPIC_API_KEY"):
|
||||
print("ERROR: ANTHROPIC_API_KEY not set", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
store = TraceStore(DB)
|
||||
# Use the TraceStore's connection for the initial read
|
||||
conn = store._conn
|
||||
conn.row_factory = sqlite3.Row
|
||||
|
||||
# Pull all traces that need scoring
|
||||
rows = list(conn.execute(
|
||||
"SELECT trace_id, query, result, agent, model FROM traces "
|
||||
"WHERE feedback IS NULL"
|
||||
))
|
||||
already_scored = conn.execute(
|
||||
"SELECT COUNT(*) FROM traces WHERE feedback IS NOT NULL"
|
||||
).fetchone()[0]
|
||||
|
||||
print(f"traces.db: {conn.execute('SELECT COUNT(*) FROM traces').fetchone()[0]} total")
|
||||
print(f" already scored: {already_scored}")
|
||||
print(f" to judge: {len(rows)}")
|
||||
print(f"parallelism: {MAX_WORKERS} workers")
|
||||
print(f"log: {LOG}")
|
||||
|
||||
if not rows:
|
||||
print("Nothing to do.")
|
||||
return 0
|
||||
|
||||
ce = CloudEngine()
|
||||
write_lock = threading.Lock()
|
||||
log_lock = threading.Lock()
|
||||
log_fp = open(LOG, "a", encoding="utf-8")
|
||||
|
||||
# Write a header line marking this run
|
||||
log_fp.write(json.dumps({
|
||||
"_run_started": datetime.utcnow().isoformat(timespec="seconds") + "Z",
|
||||
"model": MODEL,
|
||||
"workers": MAX_WORKERS,
|
||||
"to_judge": len(rows),
|
||||
"already_scored": already_scored,
|
||||
}) + "\n")
|
||||
log_fp.flush()
|
||||
|
||||
total_cost = 0.0
|
||||
done = 0
|
||||
errors = 0
|
||||
score_counts: dict = {}
|
||||
t_start = time.time()
|
||||
|
||||
def on_result(res: dict) -> None:
|
||||
nonlocal total_cost, done, errors
|
||||
trace_id = res["trace_id"]
|
||||
|
||||
# Write feedback to DB (if we got a valid score)
|
||||
if res["score"] is not None and res["error"] is None:
|
||||
with write_lock:
|
||||
store.update_feedback(trace_id, res["score"])
|
||||
conn.commit()
|
||||
else:
|
||||
errors += 1
|
||||
|
||||
with log_lock:
|
||||
log_fp.write(json.dumps(res) + "\n")
|
||||
log_fp.flush()
|
||||
|
||||
total_cost += res["cost"] or 0.0
|
||||
done += 1
|
||||
score_counts[res["score"]] = score_counts.get(res["score"], 0) + 1
|
||||
|
||||
# Progress every 50 or on error
|
||||
if done % 50 == 0 or res["error"]:
|
||||
elapsed = time.time() - t_start
|
||||
rate = done / max(0.001, elapsed)
|
||||
eta = (len(rows) - done) / max(0.001, rate)
|
||||
tag = "ERR " if res["error"] else " "
|
||||
print(
|
||||
f"{tag}[{done:4}/{len(rows)}] score={res['score']} "
|
||||
f"cost=${total_cost:.3f} "
|
||||
f"rate={rate:.1f}/s eta={eta:.0f}s "
|
||||
f"errors={errors}"
|
||||
)
|
||||
if res["error"]:
|
||||
print(f" ERROR on {trace_id[:12]}: {res['error']}")
|
||||
|
||||
with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool:
|
||||
futures = [
|
||||
pool.submit(judge_one, ce, r["trace_id"], r["query"] or "", r["result"] or "")
|
||||
for r in rows
|
||||
]
|
||||
for fut in as_completed(futures):
|
||||
res = fut.result()
|
||||
on_result(res)
|
||||
|
||||
log_fp.close()
|
||||
|
||||
elapsed = time.time() - t_start
|
||||
print(f"\n{'='*60}")
|
||||
print(f"A1 COMPLETE in {elapsed:.1f}s ({elapsed/60:.1f}min)")
|
||||
print(f"Total cost: ${total_cost:.4f}")
|
||||
print(f"Errors: {errors}/{len(rows)}")
|
||||
print(f"Score distribution:")
|
||||
for k in sorted(score_counts.keys(), key=lambda x: (x is None, x)):
|
||||
v = score_counts[k]
|
||||
pct = 100 * v / len(rows)
|
||||
label = {0.2: "failure", 0.4: "poor", 0.6: "partial", 0.8: "clean", None: "ERROR"}.get(k, "?")
|
||||
print(f" {k} ({label}): {v} ({pct:.1f}%)")
|
||||
|
||||
# Verify by re-counting from DB
|
||||
with_fb = conn.execute(
|
||||
"SELECT COUNT(*) FROM traces WHERE feedback IS NOT NULL"
|
||||
).fetchone()[0]
|
||||
above_gate = conn.execute(
|
||||
"SELECT COUNT(*) FROM traces WHERE feedback >= 0.7"
|
||||
).fetchone()[0]
|
||||
print(f"\nPost-A1 DB state:")
|
||||
print(f" traces with feedback: {with_fb}")
|
||||
print(f" traces passing 0.7 gate (eligible for personal benchmark): {above_gate}")
|
||||
return 0 if errors == 0 else 2
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,304 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Generate all distillation experiment TOML configs.
|
||||
|
||||
Produces configs for 7 experiment axes × multiple settings.
|
||||
Run: python scripts/experiments/generate_distillation_configs.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
from pathlib import Path
|
||||
|
||||
CONFIGS_DIR = Path("src/openjarvis/evals/configs/distillation")
|
||||
|
||||
# ── Teacher models ──────────────────────────────────────────────────────────
|
||||
TEACHERS = {
|
||||
"opus": {
|
||||
"model": "claude-opus-4-6",
|
||||
"engine": "cloud",
|
||||
"provider": "anthropic",
|
||||
},
|
||||
"gpt54": {
|
||||
"model": "gpt-5.4",
|
||||
"engine": "cloud",
|
||||
"provider": "openai",
|
||||
},
|
||||
"gemini": {
|
||||
"model": "gemini-3.1-pro-preview",
|
||||
"engine": "cloud",
|
||||
"provider": "google",
|
||||
},
|
||||
"qwen397b": {
|
||||
"model": "Qwen/Qwen3.5-397B-A17B-FP8",
|
||||
"engine": "vllm",
|
||||
"provider": "local",
|
||||
"note": "# Requires 8×H100, vLLM serve on port 8010",
|
||||
},
|
||||
}
|
||||
|
||||
# ── Student models ──────────────────────────────────────────────────────────
|
||||
# Served via vLLM on this H100 node. 27B uses the FP8 weights that fit on a
|
||||
# single H100; 2B and 9B use standard FP16.
|
||||
STUDENTS = {
|
||||
"2b": {"model": "Qwen/Qwen3.5-2B", "engine": "vllm", "port": 8000},
|
||||
"9b": {"model": "Qwen/Qwen3.5-9B", "engine": "vllm", "port": 8001},
|
||||
"27b": {"model": "Qwen/Qwen3.5-27B-FP8", "engine": "vllm", "port": 8002},
|
||||
}
|
||||
|
||||
# ── Benchmarks ──────────────────────────────────────────────────────────────
|
||||
BENCHMARKS = {
|
||||
"pb": "pinchbench",
|
||||
"tc15": "toolcall15",
|
||||
"tb": "taubench",
|
||||
}
|
||||
|
||||
# ── Data configs ────────────────────────────────────────────────────────────
|
||||
DATA_CONFIGS = {
|
||||
"C1": {
|
||||
"desc": "Zero test data — external traces only (GeneralThought + ADP)",
|
||||
"trace_source": "external",
|
||||
"benchmark_queries_visible": False,
|
||||
},
|
||||
"C2": {
|
||||
"desc": "Test queries only — benchmark traces visible, answers hidden",
|
||||
"trace_source": "benchmark",
|
||||
"benchmark_queries_visible": True,
|
||||
},
|
||||
"C3": {
|
||||
"desc": "Test queries + external — both benchmark and external traces",
|
||||
"trace_source": "both",
|
||||
"benchmark_queries_visible": True,
|
||||
},
|
||||
}
|
||||
|
||||
# ── Budget presets ──────────────────────────────────────────────────────────
|
||||
BUDGETS = {
|
||||
"minimal": {"max_tool_calls": 5, "max_cost": 0.50},
|
||||
"standard": {"max_tool_calls": 15, "max_cost": 2.00},
|
||||
"thorough": {"max_tool_calls": 30, "max_cost": 5.00},
|
||||
"exhaustive": {"max_tool_calls": 50, "max_cost": 10.00},
|
||||
}
|
||||
|
||||
# ── Gate presets ────────────────────────────────────────────────────────────
|
||||
GATES = {
|
||||
"permissive": {"min_improvement": 0.0, "max_regression": 0.10},
|
||||
"standard": {"min_improvement": 0.0, "max_regression": 0.05},
|
||||
"strict": {"min_improvement": 0.02, "max_regression": 0.02},
|
||||
"none": {"min_improvement": -1.0, "max_regression": 1.0},
|
||||
}
|
||||
|
||||
# ── Autonomy presets ────────────────────────────────────────────────────────
|
||||
AUTONOMY_MODES = ["auto", "tiered", "manual"]
|
||||
|
||||
|
||||
def render_config(
|
||||
*,
|
||||
experiment: str,
|
||||
teacher_key: str,
|
||||
student_key: str,
|
||||
benchmark_key: str,
|
||||
data_config_key: str = "C2",
|
||||
budget_key: str = "standard",
|
||||
gate_key: str = "standard",
|
||||
autonomy: str = "auto",
|
||||
iterative_sessions: int = 1,
|
||||
) -> str:
|
||||
"""Render a TOML config string."""
|
||||
teacher = TEACHERS[teacher_key]
|
||||
student = STUDENTS[student_key]
|
||||
benchmark = BENCHMARKS[benchmark_key]
|
||||
data_cfg = DATA_CONFIGS[data_config_key]
|
||||
budget = BUDGETS[budget_key]
|
||||
gate = GATES[gate_key]
|
||||
|
||||
note = teacher.get("note", "")
|
||||
note_line = f"\n{note}" if note else ""
|
||||
|
||||
return f"""\
|
||||
# Distillation Experiment Config
|
||||
# Experiment: {experiment}
|
||||
# Teacher: {teacher["model"]} ({teacher["provider"]})
|
||||
# Student: {student["model"]} ({student["engine"]})
|
||||
# Benchmark: {benchmark}
|
||||
# Data config: {data_config_key} — {data_cfg["desc"]}
|
||||
# Budget: {budget_key} ({budget["max_tool_calls"]} tool calls, ${budget["max_cost"]:.2f})
|
||||
# Gate: {gate_key} (min_improvement={gate["min_improvement"]}, max_regression={gate["max_regression"]})
|
||||
# Autonomy: {autonomy}
|
||||
# Iterative sessions: {iterative_sessions}
|
||||
{note_line}
|
||||
|
||||
[intelligence]
|
||||
default_model = "{student["model"]}"
|
||||
|
||||
[engine]
|
||||
default = "{student["engine"]}"
|
||||
|
||||
[engine.vllm]
|
||||
host = "http://localhost:{student.get("port", 8000)}"
|
||||
|
||||
[learning.distillation]
|
||||
enabled = true
|
||||
autonomy_mode = "{autonomy}"
|
||||
teacher_model = "{teacher["model"]}"
|
||||
max_cost_per_session_usd = {budget["max_cost"]}
|
||||
max_tool_calls_per_diagnosis = {budget["max_tool_calls"]}
|
||||
|
||||
[learning.distillation.gate]
|
||||
min_improvement = {gate["min_improvement"]}
|
||||
max_regression = {gate["max_regression"]}
|
||||
benchmark_subsample_size = 50
|
||||
|
||||
[learning.distillation.benchmark]
|
||||
synthesis_feedback_threshold = 0.7
|
||||
max_benchmark_size = 200
|
||||
|
||||
[learning.distillation.experiment]
|
||||
# Metadata for the experiment runner (not read by distillation itself)
|
||||
experiment_id = "{experiment}"
|
||||
teacher_key = "{teacher_key}"
|
||||
student_key = "{student_key}"
|
||||
benchmark = "{benchmark}"
|
||||
data_config = "{data_config_key}"
|
||||
trace_source = "{data_cfg["trace_source"]}"
|
||||
benchmark_queries_visible = {str(data_cfg["benchmark_queries_visible"]).lower()}
|
||||
budget_key = "{budget_key}"
|
||||
gate_key = "{gate_key}"
|
||||
iterative_sessions = {iterative_sessions}
|
||||
"""
|
||||
|
||||
|
||||
def write_config(subdir: str, filename: str, content: str) -> Path:
|
||||
path = CONFIGS_DIR / subdir / filename
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(content, encoding="utf-8")
|
||||
return path
|
||||
|
||||
|
||||
def generate_all() -> int:
|
||||
count = 0
|
||||
|
||||
# ── Exp 1a: Teacher Model Ablation ────────────────────────────────────
|
||||
# Fix: S-9b, B-standard, A-auto, G-standard, I-single
|
||||
# Vary: teacher × benchmark × data_config
|
||||
for teacher_key, bench_key, dc_key in itertools.product(
|
||||
TEACHERS, BENCHMARKS, DATA_CONFIGS
|
||||
):
|
||||
filename = f"{teacher_key}-9b-{bench_key}-{dc_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp1a-teacher/{teacher_key}-{bench_key}-{dc_key}",
|
||||
teacher_key=teacher_key,
|
||||
student_key="9b",
|
||||
benchmark_key=bench_key,
|
||||
data_config_key=dc_key,
|
||||
)
|
||||
write_config("exp1a-teacher", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 1b: Budget Ablation ───────────────────────────────────────────
|
||||
# Fix: S-9b, T-sonnet(opus for quality), A-auto, G-standard, I-single
|
||||
# Vary: budget × benchmark × data_config
|
||||
for budget_key, bench_key, dc_key in itertools.product(
|
||||
BUDGETS, BENCHMARKS, DATA_CONFIGS
|
||||
):
|
||||
filename = f"{budget_key}-9b-{bench_key}-{dc_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp1b-budget/{budget_key}-{bench_key}-{dc_key}",
|
||||
teacher_key="opus",
|
||||
student_key="9b",
|
||||
benchmark_key=bench_key,
|
||||
data_config_key=dc_key,
|
||||
budget_key=budget_key,
|
||||
)
|
||||
write_config("exp1b-budget", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 1c: Student Model Scaling ─────────────────────────────────────
|
||||
# Fix: T-opus, B-standard, A-auto, G-standard, I-single
|
||||
# Vary: student × benchmark × data_config
|
||||
for student_key, bench_key, dc_key in itertools.product(
|
||||
STUDENTS, BENCHMARKS, DATA_CONFIGS
|
||||
):
|
||||
filename = f"opus-{student_key}-{bench_key}-{dc_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp1c-student/opus-{student_key}-{bench_key}-{dc_key}",
|
||||
teacher_key="opus",
|
||||
student_key=student_key,
|
||||
benchmark_key=bench_key,
|
||||
data_config_key=dc_key,
|
||||
)
|
||||
write_config("exp1c-student", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 2a: Gate Strictness ───────────────────────────────────────────
|
||||
# Fix: S-9b, T-opus, B-standard, A-auto, I-single
|
||||
# Vary: gate × benchmark
|
||||
for gate_key, bench_key in itertools.product(GATES, BENCHMARKS):
|
||||
filename = f"{gate_key}-9b-{bench_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp2a-gate/{gate_key}-{bench_key}",
|
||||
teacher_key="opus",
|
||||
student_key="9b",
|
||||
benchmark_key=bench_key,
|
||||
gate_key=gate_key,
|
||||
)
|
||||
write_config("exp2a-gate", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 2b: Autonomy Mode ────────────────────────────────────────────
|
||||
# Fix: S-9b, T-opus, B-standard, G-standard, I-single
|
||||
# Vary: autonomy × benchmark
|
||||
for autonomy, bench_key in itertools.product(AUTONOMY_MODES, BENCHMARKS):
|
||||
filename = f"{autonomy}-9b-{bench_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp2b-autonomy/{autonomy}-{bench_key}",
|
||||
teacher_key="opus",
|
||||
student_key="9b",
|
||||
benchmark_key=bench_key,
|
||||
autonomy=autonomy,
|
||||
)
|
||||
write_config("exp2b-autonomy", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 3a: Iterative Sessions ───────────────────────────────────────
|
||||
# Fix: S-9b, T-opus, B-standard, A-auto, G-standard
|
||||
# Vary: number of chained sessions × benchmark
|
||||
for n_sessions, bench_key in itertools.product([1, 3, 5], BENCHMARKS):
|
||||
filename = f"iter{n_sessions}-9b-{bench_key}.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp3a-iterative/iter{n_sessions}-{bench_key}",
|
||||
teacher_key="opus",
|
||||
student_key="9b",
|
||||
benchmark_key=bench_key,
|
||||
iterative_sessions=n_sessions,
|
||||
)
|
||||
write_config("exp3a-iterative", filename, content)
|
||||
count += 1
|
||||
|
||||
# ── Exp 3b: Cross-Benchmark Transfer ─────────────────────────────────
|
||||
# Optimize using traces from benchmark X, eval on benchmark Y
|
||||
for opt_bench, eval_bench in itertools.permutations(BENCHMARKS, 2):
|
||||
filename = f"opt-{opt_bench}-eval-{eval_bench}-9b.toml"
|
||||
content = render_config(
|
||||
experiment=f"exp3b-transfer/opt-{opt_bench}-eval-{eval_bench}",
|
||||
teacher_key="opus",
|
||||
student_key="9b",
|
||||
benchmark_key=opt_bench, # Traces from this benchmark
|
||||
)
|
||||
# Add eval benchmark as metadata
|
||||
content += f'\neval_benchmark = "{BENCHMARKS[eval_bench]}"\n'
|
||||
write_config("exp3b-transfer", filename, content)
|
||||
count += 1
|
||||
|
||||
return count
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
n = generate_all()
|
||||
print(f"Generated {n} config files in {CONFIGS_DIR}/")
|
||||
|
||||
# Print summary
|
||||
for subdir in sorted(CONFIGS_DIR.iterdir()):
|
||||
if subdir.is_dir():
|
||||
files = list(subdir.glob("*.toml"))
|
||||
print(f" {subdir.name}/: {len(files)} configs")
|
||||
@@ -0,0 +1,133 @@
|
||||
#!/usr/bin/env python3
|
||||
"""M2: Collect distilled eval results and produce comparison table.
|
||||
|
||||
Reads .summary.json files from results/neurips-2026/{distilled,baselines}/
|
||||
and produces a before/after comparison against the Step 1 baseline numbers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# Jon's Step 1 baselines — used when a local baseline result doesn't exist
|
||||
STEP1_BASELINES = {
|
||||
"2b": {"toolcall15": 33.3, "pinchbench": 69.6, "livecodebench": 5.6, "taubench": 70.0, "taubench-telecom": 0.0, "gaia": 0.0, "liveresearch": 0.0, "liveresearchbench": None},
|
||||
"9b": {"toolcall15": 46.7, "pinchbench": 95.7, "livecodebench": 17.6, "taubench": 85.0, "taubench-telecom": 80.0, "gaia": 38.0, "liveresearch": 75.0, "liveresearchbench": None},
|
||||
"27b": {"toolcall15": 40.0, "pinchbench": 65.2, "livecodebench": 20.0, "taubench": 75.0, "taubench-telecom": 75.0, "gaia": 48.0, "liveresearch": 66.7, "liveresearchbench": None},
|
||||
}
|
||||
|
||||
# Which benchmarks go through the agent layer (where distillation edits actually apply)
|
||||
AGENT_BENCHMARKS = {"pinchbench", "gaia", "liveresearch"}
|
||||
|
||||
DISTILLED_ROOT = Path("results/neurips-2026/distilled")
|
||||
BASELINE_ROOT = Path("results/neurips-2026/baselines")
|
||||
|
||||
|
||||
def find_summary(root: Path, size: str, bench: str) -> Path | None:
|
||||
"""Find the summary JSON for a model × benchmark run."""
|
||||
# Expected path: root/qwen-{size}/{bench}/{bench}_Qwen-Qwen3.5-{size}.summary.json
|
||||
candidates = list(root.glob(f"qwen-{size}/{bench}/*.summary.json"))
|
||||
return candidates[0] if candidates else None
|
||||
|
||||
|
||||
def load_accuracy(summary_path: Path) -> float | None:
|
||||
"""Extract overall accuracy from a summary.json file."""
|
||||
try:
|
||||
d = json.loads(summary_path.read_text())
|
||||
# The summary has various shapes; try a few
|
||||
for key in ["overall_accuracy", "accuracy", "overall_score"]:
|
||||
if key in d:
|
||||
return float(d[key]) * 100 if d[key] <= 1.0 else float(d[key])
|
||||
# Try nested
|
||||
if "results" in d:
|
||||
for r in d["results"]:
|
||||
if "accuracy" in r:
|
||||
return float(r["accuracy"]) * 100 if r["accuracy"] <= 1.0 else float(r["accuracy"])
|
||||
except Exception as e:
|
||||
print(f" error reading {summary_path}: {e}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
|
||||
def main() -> int:
|
||||
print("=" * 100)
|
||||
print("M2 Distilled vs Baseline Comparison")
|
||||
print("=" * 100)
|
||||
print()
|
||||
print(f"{'Model':8} {'Benchmark':20} {'Baseline':>10} {'Distilled':>10} {'Delta':>10} {'Agent?':>10}")
|
||||
print("-" * 100)
|
||||
|
||||
benchmarks = [
|
||||
"toolcall15", "pinchbench", "livecodebench",
|
||||
"taubench", "taubench-telecom",
|
||||
"gaia", "liveresearch", "liveresearchbench",
|
||||
]
|
||||
|
||||
summary_rows = []
|
||||
for size in ["2b", "9b", "27b"]:
|
||||
for bench in benchmarks:
|
||||
# Baseline: prefer local file (if run this session), fall back to Step 1 numbers
|
||||
baseline_path = find_summary(BASELINE_ROOT, size, bench)
|
||||
if baseline_path:
|
||||
baseline = load_accuracy(baseline_path)
|
||||
base_source = "local"
|
||||
else:
|
||||
baseline = STEP1_BASELINES[size].get(bench)
|
||||
base_source = "step1"
|
||||
|
||||
# Distilled: must be local from this session
|
||||
distilled_path = find_summary(DISTILLED_ROOT, size, bench)
|
||||
distilled = load_accuracy(distilled_path) if distilled_path else None
|
||||
|
||||
# Format
|
||||
b_str = f"{baseline:.1f}%" if baseline is not None else "—"
|
||||
d_str = f"{distilled:.1f}%" if distilled is not None else "pending"
|
||||
if baseline is not None and distilled is not None:
|
||||
delta = distilled - baseline
|
||||
d_sign = "+" if delta >= 0 else ""
|
||||
delta_str = f"{d_sign}{delta:.1f}%"
|
||||
else:
|
||||
delta_str = "—"
|
||||
agent = "AGENT" if bench in AGENT_BENCHMARKS else "direct"
|
||||
|
||||
print(f"qwen-{size:4} {bench:20} {b_str:>10} {d_str:>10} {delta_str:>10} {agent:>10}")
|
||||
summary_rows.append({
|
||||
"model": f"qwen-{size}",
|
||||
"benchmark": bench,
|
||||
"baseline": baseline,
|
||||
"distilled": distilled,
|
||||
"delta": distilled - baseline if (baseline is not None and distilled is not None) else None,
|
||||
"agent_benchmark": bench in AGENT_BENCHMARKS,
|
||||
})
|
||||
print()
|
||||
|
||||
# Aggregate: agent vs direct benchmark deltas
|
||||
print("=" * 100)
|
||||
print("Aggregate deltas by benchmark type (paper finding)")
|
||||
print("=" * 100)
|
||||
agent_deltas = [r["delta"] for r in summary_rows if r["delta"] is not None and r["agent_benchmark"]]
|
||||
direct_deltas = [r["delta"] for r in summary_rows if r["delta"] is not None and not r["agent_benchmark"]]
|
||||
if agent_deltas:
|
||||
mean_agent = sum(agent_deltas) / len(agent_deltas)
|
||||
print(f"Agent benchmarks (PB, GAIA, DeepResearchBench): mean delta = {mean_agent:+.2f}% over {len(agent_deltas)} runs")
|
||||
if direct_deltas:
|
||||
mean_direct = sum(direct_deltas) / len(direct_deltas)
|
||||
print(f"Direct benchmarks (TC15, TauB, TBTel, LRB, LCB): mean delta = {mean_direct:+.2f}% over {len(direct_deltas)} runs")
|
||||
print()
|
||||
|
||||
# Completion progress
|
||||
expected = 24
|
||||
distilled_count = sum(1 for r in summary_rows if r["distilled"] is not None)
|
||||
print(f"Distilled runs complete: {distilled_count}/{expected}")
|
||||
|
||||
# Save JSON
|
||||
out = Path("results/neurips-2026/distillation-m2/m2_comparison.json")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
out.write_text(json.dumps(summary_rows, indent=2, default=str))
|
||||
print(f"Full data: {out}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,307 @@
|
||||
#!/usr/bin/env python3
|
||||
"""M2: Create distilled eval configs from M1 consensus edits.
|
||||
|
||||
Generates 24 distilled configs (3 models × 8 benchmarks) by cloning
|
||||
baseline configs and applying the consensus edits from M1. Also creates
|
||||
4 missing baseline configs (livecodebench-qwen-9b, liveresearchbench-*).
|
||||
|
||||
Usage: python scripts/experiments/m2_create_distilled_configs.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
CONFIGS_DIR = Path("src/openjarvis/evals/configs")
|
||||
M2_DIR = CONFIGS_DIR / "distillation" / "m2"
|
||||
|
||||
# ── Consensus values from M1 (1,131 edits) ──────────────────────────────
|
||||
DISTILLED_TEMP = 0.2 # 84/134 votes (agent benchmarks only)
|
||||
DISTILLED_MAX_TURNS = 15 # 56/125 votes (close: 25 had 49)
|
||||
REMOVE_TOOLS = {"shell_exec", "http_request"} # 13 + 6 votes
|
||||
|
||||
# ── Model specs ──────────────────────────────────────────────────────────
|
||||
MODELS = {
|
||||
"2b": {"name": "Qwen/Qwen3.5-2B", "num_gpus": 1, "port": 8000},
|
||||
"9b": {"name": "Qwen/Qwen3.5-9B", "num_gpus": 1, "port": 8001},
|
||||
"27b": {"name": "Qwen/Qwen3.5-27B-FP8", "num_gpus": 1, "port": 8002},
|
||||
}
|
||||
|
||||
# ── Benchmark specs ──────────────────────────────────────────────────────
|
||||
# Each benchmark defines its baseline config and what changes in distilled.
|
||||
BENCHMARKS = {
|
||||
"toolcall15": {
|
||||
"backend": "jarvis-direct",
|
||||
"baseline_temp": 0.0,
|
||||
"distilled_temp": 0.0, # CONTROL: no change for coding
|
||||
"max_tokens": 4096,
|
||||
"max_samples": None,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
"pinchbench": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "native_openhands",
|
||||
"baseline_temp": 0.6,
|
||||
"distilled_temp": DISTILLED_TEMP,
|
||||
"max_tokens": 8192,
|
||||
"max_samples": None,
|
||||
"judge_model": "claude-opus-4-5",
|
||||
"judge_engine": "cloud",
|
||||
"baseline_tools": [
|
||||
"think", "file_read", "file_write", "web_search", "shell_exec",
|
||||
"code_interpreter", "browser_navigate", "image_generate",
|
||||
"calculator", "http_request", "pdf_extract",
|
||||
],
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
"taubench": {
|
||||
"backend": "jarvis-direct",
|
||||
"baseline_temp": 0.7,
|
||||
"distilled_temp": DISTILLED_TEMP,
|
||||
"max_tokens": 4096,
|
||||
"max_samples": 20,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"extra_benchmark_fields": {"split": "airline,retail"},
|
||||
},
|
||||
"taubench-telecom": {
|
||||
"benchmark_name": "taubench", # same benchmark, different split
|
||||
"backend": "jarvis-direct",
|
||||
"baseline_temp": 0.7,
|
||||
"distilled_temp": DISTILLED_TEMP,
|
||||
"max_tokens": 4096,
|
||||
"max_samples": 20,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"extra_benchmark_fields": {"split": "telecom"},
|
||||
},
|
||||
"gaia": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "monitor_operative",
|
||||
"baseline_temp": 0.6,
|
||||
"distilled_temp": DISTILLED_TEMP,
|
||||
"max_tokens": 8192,
|
||||
"max_samples": 50,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"baseline_tools": [
|
||||
"think", "calculator", "code_interpreter", "web_search", "file_read",
|
||||
],
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
"liveresearch": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "monitor_operative",
|
||||
"baseline_temp": 0.6,
|
||||
"distilled_temp": DISTILLED_TEMP,
|
||||
"max_tokens": 16384,
|
||||
"max_samples": 50,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"baseline_tools": [
|
||||
"web_search", "file_read", "file_write", "code_interpreter", "think",
|
||||
],
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
"liveresearchbench": {
|
||||
"backend": "jarvis-direct",
|
||||
"baseline_temp": 0.0,
|
||||
"distilled_temp": 0.0, # CONTROL: reasoning benchmark
|
||||
"max_tokens": 8192,
|
||||
"max_samples": 50,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
"livecodebench": {
|
||||
"backend": "jarvis-direct",
|
||||
"baseline_temp": 0.0,
|
||||
"distilled_temp": 0.0, # CONTROL: coding benchmark
|
||||
"max_tokens": 4096,
|
||||
"max_samples": 20,
|
||||
"judge_model": "gpt-5-mini-2025-08-07",
|
||||
"judge_engine": "cloud",
|
||||
"extra_benchmark_fields": {},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def render_config(
|
||||
*,
|
||||
comment: str,
|
||||
meta_name: str,
|
||||
description: str,
|
||||
temperature: float,
|
||||
max_tokens: int,
|
||||
judge_model: str,
|
||||
judge_engine: str,
|
||||
output_dir: str,
|
||||
model_name: str,
|
||||
model_engine: str,
|
||||
num_gpus: int,
|
||||
benchmark_name: str,
|
||||
backend: str,
|
||||
agent: str | None = None,
|
||||
tools: list[str] | None = None,
|
||||
max_samples: int | None = None,
|
||||
extra_benchmark: dict | None = None,
|
||||
seed: int = 42,
|
||||
) -> str:
|
||||
lines = [f"# {comment}"]
|
||||
lines.append(f'[meta]\nname = "{meta_name}"\ndescription = "{description}"\n')
|
||||
lines.append(f"[defaults]\ntemperature = {temperature}\nmax_tokens = {max_tokens}\n")
|
||||
lines.append(f'[judge]\nmodel = "{judge_model}"\ntemperature = 0.0')
|
||||
if judge_engine:
|
||||
lines.append(f'engine = "{judge_engine}"')
|
||||
lines.append(f"max_tokens = 4096\n")
|
||||
lines.append(f'[run]\nmax_workers = 1\noutput_dir = "{output_dir}"\nseed = {seed}\n')
|
||||
lines.append(f'[[models]]\nname = "{model_name}"\nengine = "{model_engine}"\nnum_gpus = {num_gpus}\n')
|
||||
lines.append(f'[[benchmarks]]\nname = "{benchmark_name}"\nbackend = "{backend}"')
|
||||
if agent:
|
||||
lines.append(f'agent = "{agent}"')
|
||||
if max_samples:
|
||||
lines.append(f"max_samples = {max_samples}")
|
||||
if tools:
|
||||
tools_str = ", ".join(f'"{t}"' for t in tools)
|
||||
lines.append(f"tools = [{tools_str}]")
|
||||
if extra_benchmark:
|
||||
for k, v in extra_benchmark.items():
|
||||
if isinstance(v, str):
|
||||
lines.append(f'{k} = "{v}"')
|
||||
else:
|
||||
lines.append(f"{k} = {v}")
|
||||
lines.append("")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def make_size_label(size: str) -> str:
|
||||
return {"2b": "qwen-2b", "9b": "qwen-9b", "27b": "qwen-27b"}[size]
|
||||
|
||||
|
||||
def generate_missing_baselines() -> int:
|
||||
"""Create baseline configs that don't exist yet."""
|
||||
count = 0
|
||||
|
||||
# livecodebench-qwen-9b (missing)
|
||||
p = CONFIGS_DIR / "livecodebench-qwen-9b.toml"
|
||||
if not p.exists():
|
||||
b = BENCHMARKS["livecodebench"]
|
||||
m = MODELS["9b"]
|
||||
p.write_text(render_config(
|
||||
comment="LiveCodeBench eval: Qwen3.5-9B (vLLM, 1 GPU)",
|
||||
meta_name="livecodebench-qwen-9b",
|
||||
description="LiveCodeBench on Qwen/Qwen3.5-9B (vLLM, 1 GPU)",
|
||||
temperature=b["baseline_temp"],
|
||||
max_tokens=b["max_tokens"],
|
||||
judge_model=b["judge_model"],
|
||||
judge_engine=b["judge_engine"],
|
||||
output_dir="results/neurips-2026/baselines/qwen-9b/livecodebench/",
|
||||
model_name=m["name"], model_engine="vllm", num_gpus=m["num_gpus"],
|
||||
benchmark_name="livecodebench", backend=b["backend"],
|
||||
max_samples=b["max_samples"],
|
||||
))
|
||||
count += 1
|
||||
print(f" created {p}")
|
||||
|
||||
# liveresearchbench-qwen-{2b,9b,27b}
|
||||
for size, m in MODELS.items():
|
||||
sl = make_size_label(size)
|
||||
p = CONFIGS_DIR / f"liveresearchbench-{sl}.toml"
|
||||
if not p.exists():
|
||||
b = BENCHMARKS["liveresearchbench"]
|
||||
p.write_text(render_config(
|
||||
comment=f"LiveResearchBench (Salesforce): Qwen3.5-{size.upper()} (vLLM)",
|
||||
meta_name=f"liveresearchbench-{sl}",
|
||||
description=f"LiveResearchBench on {m['name']} (vLLM)",
|
||||
temperature=b["baseline_temp"],
|
||||
max_tokens=b["max_tokens"],
|
||||
judge_model=b["judge_model"],
|
||||
judge_engine=b["judge_engine"],
|
||||
output_dir=f"results/neurips-2026/baselines/{sl}/liveresearchbench/",
|
||||
model_name=m["name"], model_engine="vllm", num_gpus=m["num_gpus"],
|
||||
benchmark_name="liveresearchbench", backend=b["backend"],
|
||||
max_samples=b["max_samples"],
|
||||
))
|
||||
count += 1
|
||||
print(f" created {p}")
|
||||
|
||||
return count
|
||||
|
||||
|
||||
def generate_distilled_configs() -> int:
|
||||
"""Create distilled configs for all 24 model × benchmark combos."""
|
||||
M2_DIR.mkdir(parents=True, exist_ok=True)
|
||||
count = 0
|
||||
|
||||
for size, model in MODELS.items():
|
||||
sl = make_size_label(size)
|
||||
for bench_key, bench in BENCHMARKS.items():
|
||||
bench_name = bench.get("benchmark_name", bench_key)
|
||||
is_agent = bench["backend"] == "jarvis-agent"
|
||||
temp = bench["distilled_temp"]
|
||||
|
||||
# Tool list: remove broken tools for agent benchmarks
|
||||
tools = None
|
||||
if is_agent and "baseline_tools" in bench:
|
||||
tools = [t for t in bench["baseline_tools"]
|
||||
if t not in REMOVE_TOOLS]
|
||||
|
||||
fname = f"{bench_key}-{sl}-distilled.toml"
|
||||
out_path = M2_DIR / fname
|
||||
|
||||
# Determine what changed for the comment
|
||||
changes = []
|
||||
if temp != bench["baseline_temp"]:
|
||||
changes.append(f"temp {bench['baseline_temp']}→{temp}")
|
||||
if tools and set(tools) != set(bench.get("baseline_tools", [])):
|
||||
removed = set(bench.get("baseline_tools", [])) - set(tools)
|
||||
changes.append(f"removed {removed}")
|
||||
if is_agent:
|
||||
changes.append(f"max_turns 10→{DISTILLED_MAX_TURNS} (via OPENJARVIS_CONFIG)")
|
||||
change_str = "; ".join(changes) if changes else "CONTROL (no change)"
|
||||
|
||||
out_path.write_text(render_config(
|
||||
comment=f"M2 DISTILLED: {bench_key} × {model['name']} — {change_str}",
|
||||
meta_name=f"{bench_key}-{sl}-distilled",
|
||||
description=f"Distilled {bench_key} on {model['name']}",
|
||||
temperature=temp,
|
||||
max_tokens=bench["max_tokens"],
|
||||
judge_model=bench["judge_model"],
|
||||
judge_engine=bench["judge_engine"],
|
||||
output_dir=f"results/neurips-2026/distilled/{sl}/{bench_key}/",
|
||||
model_name=model["name"], model_engine="vllm", num_gpus=model["num_gpus"],
|
||||
benchmark_name=bench_name, backend=bench["backend"],
|
||||
agent=bench.get("agent"),
|
||||
tools=tools,
|
||||
max_samples=bench.get("max_samples"),
|
||||
extra_benchmark=bench.get("extra_benchmark_fields"),
|
||||
))
|
||||
count += 1
|
||||
|
||||
return count
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("=== Creating missing baseline configs ===")
|
||||
n_base = generate_missing_baselines()
|
||||
print(f"Created {n_base} missing baseline configs\n")
|
||||
|
||||
print("=== Creating distilled M2 configs ===")
|
||||
n_dist = generate_distilled_configs()
|
||||
print(f"Created {n_dist} distilled configs in {M2_DIR}/\n")
|
||||
|
||||
# Summary
|
||||
print("=== Change matrix ===")
|
||||
print(f"{'Benchmark':20} {'Backend':14} {'Temp Δ':12} {'Tool Δ':20} {'max_turns Δ':12}")
|
||||
print("-" * 80)
|
||||
for bk, b in BENCHMARKS.items():
|
||||
is_agent = b["backend"] == "jarvis-agent"
|
||||
temp_change = f"{b['baseline_temp']}→{b['distilled_temp']}" if b["baseline_temp"] != b["distilled_temp"] else "—"
|
||||
tool_change = "—"
|
||||
if is_agent and "baseline_tools" in b:
|
||||
removed = REMOVE_TOOLS & set(b.get("baseline_tools", []))
|
||||
tool_change = f"-{removed}" if removed else "—"
|
||||
mt_change = f"10→{DISTILLED_MAX_TURNS}" if is_agent else "—"
|
||||
print(f"{bk:20} {b['backend']:14} {temp_change:12} {str(tool_change):20} {mt_change:12}")
|
||||
@@ -0,0 +1,183 @@
|
||||
#!/usr/bin/env bash
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# M2: Run distilled eval configs — resumable, agent-benchmarks-first ordering
|
||||
#
|
||||
# Assumes vLLM already running: 2B:8000, 9B:8001, 27B-FP8:8002
|
||||
#
|
||||
# Usage:
|
||||
# bash m2_run_distilled_evals.sh # all
|
||||
# bash m2_run_distilled_evals.sh 9b # 9b only
|
||||
# bash m2_run_distilled_evals.sh 9b gaia # 9b + gaia only
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
set -uo pipefail
|
||||
|
||||
VENV=".venv/bin/python"
|
||||
M2_CONFIGS="src/openjarvis/evals/configs/distillation/m2"
|
||||
BASELINE_CONFIGS="src/openjarvis/evals/configs"
|
||||
M2_HOME="/scratch/user/jonsaadfalcon/openjarvis-m2"
|
||||
MODEL_FILTER=${1:-all}
|
||||
BENCH_FILTER=${2:-all}
|
||||
FORCE=${FORCE:-0} # set FORCE=1 to re-run completed configs
|
||||
|
||||
RED='\033[0;31m'; GREEN='\033[0;32m'; YELLOW='\033[1;33m'; BLUE='\033[0;34m'; NC='\033[0m'
|
||||
log() { echo -e "${BLUE}[m2]${NC} $*"; }
|
||||
ok() { echo -e "${GREEN}[ OK ]${NC} $*"; }
|
||||
warn() { echo -e "${YELLOW}[WARN]${NC} $*"; }
|
||||
fail() { echo -e "${RED}[FAIL]${NC} $*"; }
|
||||
skip() { echo -e "${YELLOW}[SKIP]${NC} $*"; }
|
||||
|
||||
declare -A MODEL_PORT=( [2b]=8000 [9b]=8001 [27b]=8002 )
|
||||
|
||||
# Agent benchmarks FIRST (where distillation impact is expected)
|
||||
AGENT_BENCHMARKS="pinchbench gaia liveresearch"
|
||||
DIRECT_BENCHMARKS="toolcall15 taubench taubench-telecom livecodebench liveresearchbench"
|
||||
ALL_BENCHMARKS="${AGENT_BENCHMARKS} ${DIRECT_BENCHMARKS}"
|
||||
|
||||
check_vllm() {
|
||||
for size in 2b 9b 27b; do
|
||||
[ "$MODEL_FILTER" != "all" ] && [ "$MODEL_FILTER" != "$size" ] && continue
|
||||
local port=${MODEL_PORT[$size]}
|
||||
if ! curl -sf "http://localhost:${port}/v1/models" >/dev/null 2>&1; then
|
||||
fail "vLLM ${size} not responding on port ${port}"
|
||||
return 1
|
||||
fi
|
||||
ok "vLLM ${size} healthy on port ${port}"
|
||||
done
|
||||
}
|
||||
|
||||
# Check if a run already completed (summary.json exists AND has a real accuracy).
|
||||
# A summary with errors=total_samples (like my earlier broken-routing tests)
|
||||
# is treated as incomplete and re-run.
|
||||
is_complete() {
|
||||
local summary_path=$1
|
||||
[ -f "$summary_path" ] || return 1
|
||||
python3 -c "
|
||||
import json, sys
|
||||
try:
|
||||
d = json.load(open('$summary_path'))
|
||||
total = d.get('total_samples', 0)
|
||||
errors = d.get('errors', 0)
|
||||
scored = d.get('scored_samples', 0)
|
||||
# Complete if at least some samples were scored successfully
|
||||
sys.exit(0 if scored > 0 else 1)
|
||||
except Exception:
|
||||
sys.exit(1)
|
||||
" 2>/dev/null
|
||||
}
|
||||
|
||||
run_eval() {
|
||||
local config_path=$1 label=$2 size=$3 summary_path=$4 use_distilled=${5:-false}
|
||||
|
||||
if [ "$FORCE" != "1" ] && is_complete "$summary_path"; then
|
||||
skip "${label} [already complete]"
|
||||
return 0
|
||||
fi
|
||||
|
||||
local oj_config="${M2_HOME}/config-baseline-${size}.toml"
|
||||
[ "$use_distilled" = "true" ] && oj_config="${M2_HOME}/config-${size}.toml"
|
||||
|
||||
log "Running: ${label} [$(basename ${oj_config})]"
|
||||
env OPENJARVIS_CONFIG="${oj_config}" ${VENV} -m openjarvis.evals run -c "${config_path}" 2>&1
|
||||
local rc=$?
|
||||
if [ $rc -eq 0 ] && is_complete "$summary_path"; then
|
||||
ok "Done: ${label}"
|
||||
else
|
||||
warn "Failed: ${label} (rc=$rc)"
|
||||
fi
|
||||
}
|
||||
|
||||
# Derive the expected summary.json path for a given distilled config
|
||||
summary_for_distilled() {
|
||||
local bench=$1 size=$2
|
||||
# Output dir from the config template: results/neurips-2026/distilled/qwen-{size}/{bench}/
|
||||
# Summary file pattern: {bench}_Qwen-Qwen3.5-{size}.summary.json
|
||||
local model_slug
|
||||
if [ "$size" = "27b" ]; then model_slug="Qwen-Qwen3.5-27B-FP8"
|
||||
else model_slug="Qwen-Qwen3.5-${size}"; fi
|
||||
# Uppercase B for the model slug
|
||||
model_slug=$(echo "$model_slug" | sed 's/-\([0-9][0-9]*\)b/-\1B/g')
|
||||
# For taubench-telecom, the benchmark name in output is "taubench" not "taubench-telecom"
|
||||
local bench_fname=$bench
|
||||
[ "$bench" = "taubench-telecom" ] && bench_fname="taubench"
|
||||
echo "results/neurips-2026/distilled/qwen-${size}/${bench}/${bench_fname}_${model_slug}.summary.json"
|
||||
}
|
||||
|
||||
summary_for_baseline() {
|
||||
local bench=$1 size=$2
|
||||
local model_slug
|
||||
if [ "$size" = "27b" ]; then model_slug="Qwen-Qwen3.5-27B-FP8"
|
||||
else model_slug="Qwen-Qwen3.5-${size}"; fi
|
||||
model_slug=$(echo "$model_slug" | sed 's/-\([0-9][0-9]*\)b/-\1B/g')
|
||||
local bench_fname=$bench
|
||||
[ "$bench" = "taubench-telecom" ] && bench_fname="taubench"
|
||||
echo "results/neurips-2026/baselines/qwen-${size}/${bench}/${bench_fname}_${model_slug}.summary.json"
|
||||
}
|
||||
|
||||
log "M2 Distilled Eval Runner (resumable, agent-first)"
|
||||
log "Model filter: ${MODEL_FILTER} Benchmark filter: ${BENCH_FILTER} FORCE=${FORCE}"
|
||||
|
||||
check_vllm || exit 1
|
||||
|
||||
start_time=$(date +%s)
|
||||
|
||||
# Phase B1: DISTILLED agent benchmarks (highest-priority: PinchBench, GAIA, DeepResearchBench)
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
log "Phase B1: DISTILLED agent benchmarks (9 runs — the critical data)"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
for size in 2b 9b 27b; do
|
||||
[ "$MODEL_FILTER" != "all" ] && [ "$MODEL_FILTER" != "$size" ] && continue
|
||||
for bench in ${AGENT_BENCHMARKS}; do
|
||||
[ "$BENCH_FILTER" != "all" ] && [ "$BENCH_FILTER" != "$bench" ] && continue
|
||||
cfg="${M2_CONFIGS}/${bench}-qwen-${size}-distilled.toml"
|
||||
sum=$(summary_for_distilled "$bench" "$size")
|
||||
[ -f "$cfg" ] && run_eval "$cfg" "DISTILLED ${bench}-qwen-${size}" "${size}" "$sum" true
|
||||
done
|
||||
done
|
||||
|
||||
# Phase B2: DISTILLED direct benchmarks (controls — should show minimal delta)
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
log "Phase B2: DISTILLED direct benchmarks (15 runs — controls)"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
for size in 2b 9b 27b; do
|
||||
[ "$MODEL_FILTER" != "all" ] && [ "$MODEL_FILTER" != "$size" ] && continue
|
||||
for bench in ${DIRECT_BENCHMARKS}; do
|
||||
[ "$BENCH_FILTER" != "all" ] && [ "$BENCH_FILTER" != "$bench" ] && continue
|
||||
cfg="${M2_CONFIGS}/${bench}-qwen-${size}-distilled.toml"
|
||||
sum=$(summary_for_distilled "$bench" "$size")
|
||||
[ -f "$cfg" ] && run_eval "$cfg" "DISTILLED ${bench}-qwen-${size}" "${size}" "$sum" true
|
||||
done
|
||||
done
|
||||
|
||||
# Phase A: LiveResearchBench baselines (last, since Step 1 baselines exist for other benchmarks)
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
log "Phase A: LiveResearchBench baselines (3 runs — new benchmark only)"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
for size in 2b 9b 27b; do
|
||||
[ "$MODEL_FILTER" != "all" ] && [ "$MODEL_FILTER" != "$size" ] && continue
|
||||
[ "$BENCH_FILTER" != "all" ] && [ "$BENCH_FILTER" != "liveresearchbench" ] && continue
|
||||
cfg="${BASELINE_CONFIGS}/liveresearchbench-qwen-${size}.toml"
|
||||
sum=$(summary_for_baseline "liveresearchbench" "$size")
|
||||
[ -f "$cfg" ] && run_eval "$cfg" "BASELINE liveresearchbench-qwen-${size}" "${size}" "$sum" false
|
||||
done
|
||||
|
||||
# Phase C: Spot-check 2 baselines against Jon's Step 1 numbers
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
log "Phase C: Spot-check baselines (TC15-9B, GAIA-9B)"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
if [ "$MODEL_FILTER" = "all" ] || [ "$MODEL_FILTER" = "9b" ]; then
|
||||
for bench in toolcall15 gaia; do
|
||||
[ "$BENCH_FILTER" != "all" ] && [ "$BENCH_FILTER" != "$bench" ] && continue
|
||||
cfg="${BASELINE_CONFIGS}/${bench}-qwen-9b.toml"
|
||||
sum=$(summary_for_baseline "$bench" "9b")
|
||||
[ -f "$cfg" ] && run_eval "$cfg" "SPOTCHECK ${bench}-qwen-9b" "9b" "$sum" false
|
||||
done
|
||||
fi
|
||||
|
||||
end_time=$(date +%s)
|
||||
elapsed=$((end_time - start_time))
|
||||
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
ok "M2 complete in ${elapsed}s ($(( elapsed / 3600 ))h $(( (elapsed % 3600) / 60 ))m)"
|
||||
log "Distilled results: results/neurips-2026/distilled/"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
@@ -0,0 +1,581 @@
|
||||
#!/usr/bin/env python3
|
||||
"""M3: Empirical hill-climbing with an LLM proposer.
|
||||
|
||||
Replaces M1's open-loop "aggregate consensus across sessions" with a closed
|
||||
loop:
|
||||
|
||||
For each target (student, benchmark, agent):
|
||||
for round in 1..N:
|
||||
edit = teacher.propose_one(history_with_measured_deltas)
|
||||
score_new = eval_subsample(apply(edit))
|
||||
if score_new > current_score: accept
|
||||
|
||||
Every proposal is empirically verified before the next is proposed, and the
|
||||
teacher sees measured deltas (not just traces) in its context.
|
||||
|
||||
Usage:
|
||||
python scripts/experiments/m3_hill_climb.py \\
|
||||
--student 9b --benchmark liveresearch \\
|
||||
--rounds 4 --k-subsample 8 --k-final 30
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.core.types import Message, Role
|
||||
from openjarvis.engine.cloud import CloudEngine
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Config & constants
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
STUDENT = {
|
||||
"2b": {"name": "Qwen/Qwen3.5-2B", "port": 8000, "gpu": 4},
|
||||
"9b": {"name": "Qwen/Qwen3.5-9B", "port": 8001, "gpu": 5},
|
||||
"27b": {"name": "Qwen/Qwen3.5-27B-FP8", "port": 8002, "gpu": 6},
|
||||
}
|
||||
|
||||
# Per-benchmark defaults (backend, baseline config)
|
||||
BENCHMARK = {
|
||||
"liveresearch": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "monitor_operative",
|
||||
"baseline_temp": 0.6,
|
||||
"baseline_max_tokens": 16384,
|
||||
"baseline_max_turns": 10,
|
||||
"baseline_tools": ["web_search", "file_read", "file_write",
|
||||
"code_interpreter", "think"],
|
||||
"max_samples_final": 50, # final eval
|
||||
"judge": "gpt-5-mini-2025-08-07",
|
||||
},
|
||||
"gaia": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "monitor_operative",
|
||||
"baseline_temp": 0.6,
|
||||
"baseline_max_tokens": 8192,
|
||||
"baseline_max_turns": 10,
|
||||
"baseline_tools": ["think", "calculator", "code_interpreter",
|
||||
"web_search", "file_read"],
|
||||
"max_samples_final": 50,
|
||||
"judge": "gpt-5-mini-2025-08-07",
|
||||
},
|
||||
"pinchbench": {
|
||||
"backend": "jarvis-agent",
|
||||
"agent": "native_openhands",
|
||||
"baseline_temp": 0.6,
|
||||
"baseline_max_tokens": 8192,
|
||||
"baseline_max_turns": 10,
|
||||
"baseline_tools": ["think", "file_read", "file_write", "web_search",
|
||||
"shell_exec", "code_interpreter", "browser_navigate",
|
||||
"image_generate", "calculator", "http_request",
|
||||
"pdf_extract"],
|
||||
"max_samples_final": 23,
|
||||
"judge": "claude-opus-4-5",
|
||||
},
|
||||
}
|
||||
|
||||
AVAILABLE_TOOLS = [
|
||||
"think", "file_read", "file_write", "web_search", "shell_exec",
|
||||
"code_interpreter", "browser_navigate", "image_generate", "calculator",
|
||||
"http_request", "pdf_extract", "pdf_reader", "list_directory",
|
||||
]
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Proposer (LLM)
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
PROPOSER_SYSTEM = """\
|
||||
You are optimizing an OpenJarvis agent configuration for maximum accuracy on \
|
||||
a benchmark. You propose ONE config edit per round. After each proposal, the \
|
||||
edit is applied and the benchmark is run on a subsample; you then see the \
|
||||
measured score delta and decide the next edit.
|
||||
|
||||
Your job is to find the config that maximizes measured accuracy.
|
||||
|
||||
EDIT GRAMMAR — return JSON with "op" and the parameter fields at top level:
|
||||
|
||||
{"op": "<op_name>", <param_fields>, "rationale": "<one sentence>"}
|
||||
|
||||
Valid ops and their parameter fields:
|
||||
|
||||
1. set_temperature: "value" (float, 0.0..1.0)
|
||||
2. set_max_turns: "value" (int, 1..100)
|
||||
3. set_max_tokens: "value" (int, 512..32768)
|
||||
4. add_tool: "tool_name" (string; must be in AVAILABLE_TOOLS, not already active)
|
||||
5. remove_tool: "tool_name" (string; must be in current tools)
|
||||
6. noop: (no params; propose only if you believe no further edit will help)
|
||||
|
||||
CONCRETE EXAMPLES:
|
||||
{"op": "set_temperature", "value": 0.3, "rationale": "Reduce loop risk."}
|
||||
{"op": "set_max_turns", "value": 20, "rationale": "More turns for research tasks."}
|
||||
{"op": "add_tool", "tool_name": "pdf_extract", "rationale": "Tasks require PDF reading."}
|
||||
{"op": "remove_tool", "tool_name": "shell_exec", "rationale": "Tool is broken in this env."}
|
||||
{"op": "noop", "rationale": "Current config seems optimal."}
|
||||
|
||||
EXPLORATION BIAS:
|
||||
The config space has 5 distinct axes: temperature, max_turns, max_tokens,
|
||||
tool additions (add_tool), tool removals (remove_tool). Before proposing a
|
||||
second edit on an axis you've already tried, consider whether an untried
|
||||
axis might reveal a larger gain. Tool-list edits (add/remove) often matter
|
||||
more than numeric hyperparameters on benchmarks where the agent uses tools.
|
||||
|
||||
Return ONLY the JSON object, no preamble, no code fences."""
|
||||
|
||||
|
||||
def build_user_prompt(
|
||||
*, benchmark: str, student: str, agent: str,
|
||||
baseline_score: float, current_score: float, current_config: dict,
|
||||
edit_history: list[dict], available_tools: list[str],
|
||||
sample_queries: list[str],
|
||||
) -> str:
|
||||
hist_lines = []
|
||||
for i, e in enumerate(edit_history, 1):
|
||||
delta = e["score_after"] - e["score_before"]
|
||||
status = "ACCEPTED" if e["accepted"] else "REJECTED"
|
||||
hist_lines.append(
|
||||
f" Round {i}: {json.dumps(e['edit'])} "
|
||||
f"→ score {e['score_before']:.1f}% → {e['score_after']:.1f}% "
|
||||
f"(Δ {delta:+.1f}, {status})"
|
||||
)
|
||||
hist = "\n".join(hist_lines) if hist_lines else " (no edits yet — baseline is the starting point)"
|
||||
|
||||
samples = "\n".join(f" - {q[:150]}..." for q in sample_queries[:3])
|
||||
|
||||
unused_tools = [t for t in available_tools if t not in current_config["tools"]]
|
||||
|
||||
return f"""\
|
||||
TARGET:
|
||||
student: {student} (vLLM-served Qwen3.5)
|
||||
benchmark: {benchmark}
|
||||
agent: {agent} (uses OpenAI-format structured tool calls)
|
||||
|
||||
CURRENT CONFIG:
|
||||
temperature = {current_config['temperature']}
|
||||
max_turns = {current_config['max_turns']}
|
||||
max_tokens = {current_config['max_tokens']}
|
||||
tools = {current_config['tools']}
|
||||
|
||||
TOOLS NOT CURRENTLY ACTIVE (available to add):
|
||||
{unused_tools}
|
||||
|
||||
BASELINE (unedited) SCORE: {baseline_score:.1f}%
|
||||
CURRENT BEST SCORE: {current_score:.1f}%
|
||||
|
||||
EDIT HISTORY (with measured deltas):
|
||||
{hist}
|
||||
|
||||
SAMPLE TASKS FROM THIS BENCHMARK:
|
||||
{samples}
|
||||
|
||||
Propose ONE edit that you predict will improve the measured accuracy. Consider \
|
||||
the edit history — do not repeat proposals that were rejected. If you believe \
|
||||
further edits will not help, propose noop.
|
||||
|
||||
Return JSON only."""
|
||||
|
||||
|
||||
def call_proposer(
|
||||
engine: CloudEngine, system: str, user: str, model: str = "claude-sonnet-4-6"
|
||||
) -> dict:
|
||||
resp = engine.generate(
|
||||
messages=[
|
||||
Message(role=Role.SYSTEM, content=system),
|
||||
Message(role=Role.USER, content=user),
|
||||
],
|
||||
model=model, max_tokens=600, temperature=0.3,
|
||||
)
|
||||
content = (resp.get("content") or "").strip()
|
||||
# Be forgiving about code fences
|
||||
if content.startswith("```"):
|
||||
content = re.sub(r"^```(?:json)?\s*", "", content)
|
||||
content = re.sub(r"\s*```\s*$", "", content)
|
||||
# Find the JSON object
|
||||
m = re.search(r"\{[\s\S]*\}", content)
|
||||
if not m:
|
||||
raise ValueError(f"No JSON object found in proposer output: {content[:200]}")
|
||||
return json.loads(m.group(0))
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Edit application & evaluation
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
@dataclass
|
||||
class Config:
|
||||
temperature: float
|
||||
max_turns: int
|
||||
max_tokens: int
|
||||
tools: list[str]
|
||||
|
||||
|
||||
def apply_edit(cfg: Config, edit: dict) -> Config:
|
||||
"""Apply an edit. Tolerant of both flat and nested (params) forms."""
|
||||
op = edit["op"]
|
||||
# Merge top-level edit fields with params for flat-or-nested tolerance
|
||||
p = {**edit.get("params", {}), **{k: v for k, v in edit.items()
|
||||
if k not in ("op", "params", "rationale")}}
|
||||
new = Config(
|
||||
temperature=cfg.temperature, max_turns=cfg.max_turns,
|
||||
max_tokens=cfg.max_tokens, tools=list(cfg.tools),
|
||||
)
|
||||
if op == "noop":
|
||||
return new
|
||||
if op == "set_temperature":
|
||||
new.temperature = float(p["value"])
|
||||
elif op == "set_max_turns":
|
||||
new.max_turns = int(p["value"])
|
||||
elif op == "set_max_tokens":
|
||||
new.max_tokens = int(p["value"])
|
||||
elif op == "add_tool":
|
||||
tool = p["tool_name"]
|
||||
if tool not in new.tools:
|
||||
new.tools.append(tool)
|
||||
elif op == "remove_tool":
|
||||
tool = p["tool_name"]
|
||||
new.tools = [t for t in new.tools if t != tool]
|
||||
else:
|
||||
raise ValueError(f"unknown edit op: {op}")
|
||||
return new
|
||||
|
||||
|
||||
def write_eval_toml(
|
||||
*, bench: str, bench_spec: dict, student: dict, cfg: Config,
|
||||
k_samples: int, output_dir: Path,
|
||||
) -> Path:
|
||||
agent_line = f'agent = "{bench_spec["agent"]}"' if bench_spec.get("agent") else ""
|
||||
tools_str = "[" + ", ".join(f'"{t}"' for t in cfg.tools) + "]"
|
||||
toml = f"""\
|
||||
[meta]
|
||||
name = "m3-{bench}-{student['name'].replace('/', '-')}"
|
||||
description = "M3 hill-climb round"
|
||||
|
||||
[defaults]
|
||||
temperature = {cfg.temperature}
|
||||
max_tokens = {cfg.max_tokens}
|
||||
|
||||
[judge]
|
||||
model = "{bench_spec['judge']}"
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "{output_dir}"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "{student['name']}"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "{bench}"
|
||||
backend = "{bench_spec['backend']}"
|
||||
{agent_line}
|
||||
max_samples = {k_samples}
|
||||
tools = {tools_str}
|
||||
"""
|
||||
path = output_dir / "eval.toml"
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(toml)
|
||||
return path
|
||||
|
||||
|
||||
def write_openjarvis_config(home: Path, port: int, max_turns: int) -> Path:
|
||||
p = home / "global-config.toml"
|
||||
p.parent.mkdir(parents=True, exist_ok=True)
|
||||
p.write_text(f"""\
|
||||
[agent]
|
||||
max_turns = {max_turns}
|
||||
|
||||
[engine]
|
||||
default = "vllm"
|
||||
|
||||
[engine.vllm]
|
||||
host = "http://localhost:{port}"
|
||||
""")
|
||||
return p
|
||||
|
||||
|
||||
def run_eval(eval_toml: Path, oj_config: Path) -> tuple[float, int, int]:
|
||||
"""Run one eval. Returns (accuracy_pct, scored, total)."""
|
||||
env = {**os.environ, "OPENJARVIS_CONFIG": str(oj_config)}
|
||||
result = subprocess.run(
|
||||
[".venv/bin/python", "-m", "openjarvis.evals", "run", "-c", str(eval_toml)],
|
||||
capture_output=True, text=True, env=env, timeout=7200,
|
||||
)
|
||||
# Find the summary.json
|
||||
out_dir = eval_toml.parent
|
||||
sums = list(out_dir.glob("**/*.summary.json"))
|
||||
if not sums:
|
||||
print(f"[m3] WARNING: no summary.json in {out_dir}")
|
||||
print(f"[m3] stderr tail: {result.stderr[-500:]}")
|
||||
return 0.0, 0, 0
|
||||
d = json.loads(sums[0].read_text())
|
||||
acc = d.get("accuracy", 0.0)
|
||||
acc_pct = acc * 100 if acc <= 1.0 else acc
|
||||
return acc_pct, d.get("scored_samples", 0), d.get("total_samples", 0)
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Benchmark sample loader (for proposer context)
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def load_sample_queries(bench: str, n: int = 3) -> list[str]:
|
||||
try:
|
||||
if bench == "liveresearch":
|
||||
from openjarvis.evals.datasets.liveresearch import LiveResearchBenchDataset
|
||||
ds = LiveResearchBenchDataset()
|
||||
ds.load(max_samples=n)
|
||||
elif bench == "pinchbench":
|
||||
from openjarvis.evals.datasets.pinchbench import PinchBenchDataset
|
||||
ds = PinchBenchDataset()
|
||||
elif bench == "gaia":
|
||||
from openjarvis.evals.datasets.gaia import GAIADataset
|
||||
ds = GAIADataset()
|
||||
if hasattr(ds, "load"):
|
||||
ds.load(max_samples=n)
|
||||
else:
|
||||
return []
|
||||
return [r.problem for r in list(ds.iter_records())[:n]]
|
||||
except Exception as e:
|
||||
print(f"[m3] WARNING: could not load samples for {bench}: {e}")
|
||||
return []
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
# Main hill-climb loop
|
||||
# ═══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def hill_climb(args) -> dict:
|
||||
bench_spec = BENCHMARK[args.benchmark]
|
||||
student = STUDENT[args.student]
|
||||
|
||||
# State dir (resumable)
|
||||
base_dir = Path(args.out_dir) / f"{args.student}-{args.benchmark}"
|
||||
base_dir.mkdir(parents=True, exist_ok=True)
|
||||
state_path = base_dir / "state.json"
|
||||
|
||||
# Load or initialize state
|
||||
if state_path.exists() and not args.fresh:
|
||||
state = json.loads(state_path.read_text())
|
||||
print(f"[m3] Resumed state from round {len(state['history'])}/{args.rounds}")
|
||||
else:
|
||||
baseline_cfg_dict = {
|
||||
"temperature": bench_spec["baseline_temp"],
|
||||
"max_turns": bench_spec["baseline_max_turns"],
|
||||
"max_tokens": bench_spec["baseline_max_tokens"],
|
||||
"tools": list(bench_spec["baseline_tools"]),
|
||||
}
|
||||
# ALWAYS measure the baseline today first (unless user provides --trust-baseline).
|
||||
# This prevents anchoring to an unreproducible Step 1 number.
|
||||
# We measure at k=k_final so the final delta is like-for-like.
|
||||
if args.trust_baseline and args.baseline_score is not None:
|
||||
measured_baseline = args.baseline_score
|
||||
measured_baseline_k = None
|
||||
print(f"[m3] Trusting provided baseline score: {measured_baseline:.1f}% "
|
||||
f"(--trust-baseline set)")
|
||||
else:
|
||||
print(f"[m3] Measuring today's baseline with k={args.k_final} "
|
||||
f"(matches k_final for clean like-for-like delta)...")
|
||||
bl_dir = base_dir / "baseline_measure"
|
||||
bl_toml = write_eval_toml(
|
||||
bench=args.benchmark, bench_spec=bench_spec, student=student,
|
||||
cfg=Config(**baseline_cfg_dict),
|
||||
k_samples=args.k_final, output_dir=bl_dir,
|
||||
)
|
||||
bl_oj = write_openjarvis_config(bl_dir, student["port"], baseline_cfg_dict["max_turns"])
|
||||
t0 = time.monotonic()
|
||||
measured_baseline, bl_scored, bl_total = run_eval(bl_toml, bl_oj)
|
||||
measured_baseline_k = bl_scored
|
||||
print(f"[m3] Measured baseline: {measured_baseline:.1f}% "
|
||||
f"({bl_scored}/{bl_total}) in {(time.monotonic() - t0)/60:.1f} min")
|
||||
if args.baseline_score is not None:
|
||||
print(f"[m3] (reference: --baseline-score was {args.baseline_score:.1f}%, "
|
||||
f"drift = {measured_baseline - args.baseline_score:+.1f})")
|
||||
|
||||
state = {
|
||||
"args": vars(args),
|
||||
"benchmark": args.benchmark,
|
||||
"student": student["name"],
|
||||
"agent": bench_spec["agent"],
|
||||
"baseline_config": baseline_cfg_dict,
|
||||
"baseline_score": measured_baseline,
|
||||
"baseline_score_reference": args.baseline_score,
|
||||
"current_config": dict(baseline_cfg_dict),
|
||||
"current_score": measured_baseline,
|
||||
"history": [],
|
||||
}
|
||||
|
||||
# Save state helper
|
||||
def save():
|
||||
state_path.write_text(json.dumps(state, indent=2, default=str))
|
||||
|
||||
save()
|
||||
|
||||
# Sample queries for proposer context
|
||||
sample_queries = load_sample_queries(args.benchmark)
|
||||
|
||||
engine = CloudEngine()
|
||||
|
||||
# Hill-climb rounds
|
||||
for round_num in range(len(state["history"]) + 1, args.rounds + 1):
|
||||
print(f"\n{'═' * 70}")
|
||||
print(f"[m3] Round {round_num}/{args.rounds}")
|
||||
print(f"[m3] current_score = {state['current_score']:.1f}%")
|
||||
print(f"[m3] current_config = {state['current_config']}")
|
||||
|
||||
# Propose
|
||||
user_prompt = build_user_prompt(
|
||||
benchmark=args.benchmark, student=student["name"],
|
||||
agent=bench_spec["agent"],
|
||||
baseline_score=state["baseline_score"],
|
||||
current_score=state["current_score"],
|
||||
current_config=state["current_config"],
|
||||
edit_history=state["history"],
|
||||
available_tools=AVAILABLE_TOOLS,
|
||||
sample_queries=sample_queries,
|
||||
)
|
||||
try:
|
||||
edit = call_proposer(engine, PROPOSER_SYSTEM, user_prompt,
|
||||
model=args.proposer_model)
|
||||
except Exception as e:
|
||||
print(f"[m3] Proposer failed: {e}. Ending hill-climb.")
|
||||
break
|
||||
|
||||
print(f"[m3] Proposed: {json.dumps(edit)}")
|
||||
|
||||
if edit.get("op") == "noop":
|
||||
print(f"[m3] Teacher proposed noop; stopping.")
|
||||
break
|
||||
|
||||
# Apply and evaluate subsample
|
||||
try:
|
||||
candidate = apply_edit(Config(**state["current_config"]), edit)
|
||||
except Exception as e:
|
||||
print(f"[m3] apply_edit failed: {e}. Recording as malformed edit.")
|
||||
# Record as a rejected malformed edit so teacher won't repeat
|
||||
state["history"].append({
|
||||
"round": round_num, "edit": edit,
|
||||
"config_after": None,
|
||||
"score_before": state["current_score"],
|
||||
"score_after": state["current_score"], # no change
|
||||
"scored": 0, "total": 0, "elapsed_seconds": 0,
|
||||
"accepted": False,
|
||||
"error": f"malformed_edit: {e}",
|
||||
})
|
||||
save()
|
||||
continue
|
||||
|
||||
round_dir = base_dir / f"round_{round_num}"
|
||||
eval_toml = write_eval_toml(
|
||||
bench=args.benchmark, bench_spec=bench_spec, student=student,
|
||||
cfg=candidate, k_samples=args.k_subsample,
|
||||
output_dir=round_dir,
|
||||
)
|
||||
oj_cfg = write_openjarvis_config(round_dir, student["port"], candidate.max_turns)
|
||||
|
||||
print(f"[m3] Running k={args.k_subsample} subsample eval...")
|
||||
t0 = time.monotonic()
|
||||
acc, scored, total = run_eval(eval_toml, oj_cfg)
|
||||
elapsed = time.monotonic() - t0
|
||||
print(f"[m3] Subsample score: {acc:.1f}% ({scored}/{total}) in {elapsed/60:.1f} min")
|
||||
|
||||
score_before = state["current_score"]
|
||||
delta = acc - score_before
|
||||
accepted = delta > args.accept_threshold
|
||||
|
||||
# Record history
|
||||
state["history"].append({
|
||||
"round": round_num, "edit": edit,
|
||||
"config_after": asdict(candidate),
|
||||
"score_before": score_before, "score_after": acc,
|
||||
"scored": scored, "total": total, "elapsed_seconds": elapsed,
|
||||
"accepted": accepted,
|
||||
})
|
||||
|
||||
if accepted:
|
||||
state["current_config"] = asdict(candidate)
|
||||
state["current_score"] = acc
|
||||
print(f"[m3] ACCEPTED (Δ={delta:+.1f})")
|
||||
else:
|
||||
print(f"[m3] REJECTED (Δ={delta:+.1f} ≤ {args.accept_threshold})")
|
||||
|
||||
save()
|
||||
|
||||
# Final eval with current config
|
||||
print(f"\n{'═' * 70}")
|
||||
print(f"[m3] Final eval with best config: {state['current_config']}")
|
||||
final_dir = base_dir / "final"
|
||||
final_cfg = Config(**state["current_config"])
|
||||
eval_toml = write_eval_toml(
|
||||
bench=args.benchmark, bench_spec=bench_spec, student=student,
|
||||
cfg=final_cfg, k_samples=args.k_final,
|
||||
output_dir=final_dir,
|
||||
)
|
||||
oj_cfg = write_openjarvis_config(final_dir, student["port"], final_cfg.max_turns)
|
||||
|
||||
t0 = time.monotonic()
|
||||
final_acc, final_scored, final_total = run_eval(eval_toml, oj_cfg)
|
||||
elapsed = time.monotonic() - t0
|
||||
|
||||
state["final_score"] = final_acc
|
||||
state["final_scored"] = final_scored
|
||||
state["final_total"] = final_total
|
||||
state["final_elapsed_seconds"] = elapsed
|
||||
save()
|
||||
|
||||
print(f"\n{'═' * 70}")
|
||||
print(f"[m3] DONE")
|
||||
print(f"[m3] baseline = {state['baseline_score']:.1f}%")
|
||||
print(f"[m3] final = {final_acc:.1f}% ({final_scored}/{final_total}) in {elapsed/60:.1f} min")
|
||||
print(f"[m3] Δ vs baseline = {final_acc - state['baseline_score']:+.1f}")
|
||||
print(f"[m3] state: {state_path}")
|
||||
|
||||
return state
|
||||
|
||||
|
||||
def main() -> int:
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--student", required=True, choices=list(STUDENT))
|
||||
ap.add_argument("--benchmark", required=True, choices=list(BENCHMARK))
|
||||
ap.add_argument("--rounds", type=int, default=4)
|
||||
ap.add_argument("--k-subsample", type=int, default=8)
|
||||
ap.add_argument("--k-final", type=int, default=30)
|
||||
ap.add_argument("--baseline-score", type=float, default=None,
|
||||
help="Optional reference baseline (shown alongside measured). "
|
||||
"Hill-climb always measures today's baseline unless --trust-baseline.")
|
||||
ap.add_argument("--trust-baseline", action="store_true",
|
||||
help="Skip baseline re-measurement, trust --baseline-score.")
|
||||
ap.add_argument("--accept-threshold", type=float, default=0.0,
|
||||
help="Accept edit if score Δ > this (default: 0)")
|
||||
ap.add_argument("--proposer-model", default="claude-sonnet-4-6")
|
||||
ap.add_argument("--out-dir", default="results/neurips-2026/distillation-m3")
|
||||
ap.add_argument("--fresh", action="store_true",
|
||||
help="Overwrite existing state and start fresh")
|
||||
args = ap.parse_args()
|
||||
|
||||
result = hill_climb(args)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
+307
@@ -0,0 +1,307 @@
|
||||
#!/usr/bin/env bash
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
# Run distillation ablation experiments
|
||||
#
|
||||
# Prerequisites:
|
||||
# - Ollama running with qwen3.5:{2b,9b,27b}
|
||||
# - ANTHROPIC_API_KEY set (for Opus teacher)
|
||||
# - OPENAI_API_KEY set (for GPT-5.4 teacher)
|
||||
# - GOOGLE_API_KEY set (for Gemini teacher)
|
||||
# - For Qwen-397B teacher: vLLM serving on port 8010 with 8×H100
|
||||
# - Traces seeded with feedback (run A1 blocker first)
|
||||
# - jarvis learning init already run
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/experiments/run_distillation_experiments.sh # Run all
|
||||
# bash scripts/experiments/run_distillation_experiments.sh exp1a # Run Phase 1a only
|
||||
# bash scripts/experiments/run_distillation_experiments.sh exp1a opus # Single config
|
||||
# ──────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
CONFIGS_DIR="src/openjarvis/evals/configs/distillation"
|
||||
RESULTS_DIR="results/neurips-2026/agent-optimization/distillation"
|
||||
EXPERIMENT=${1:-all}
|
||||
FILTER=${2:-}
|
||||
|
||||
# ── Colors ───────────────────────────────────────────────────────────────────
|
||||
RED='\033[0;31m'
|
||||
GREEN='\033[0;32m'
|
||||
YELLOW='\033[1;33m'
|
||||
BLUE='\033[0;34m'
|
||||
NC='\033[0m'
|
||||
|
||||
log() { echo -e "${BLUE}[distill]${NC} $*"; }
|
||||
ok() { echo -e "${GREEN}[ OK ]${NC} $*"; }
|
||||
warn() { echo -e "${YELLOW}[ WARN ]${NC} $*"; }
|
||||
fail() { echo -e "${RED}[ FAIL ]${NC} $*"; }
|
||||
|
||||
# ── Preflight checks ────────────────────────────────────────────────────────
|
||||
check_prereqs() {
|
||||
log "Preflight checks..."
|
||||
|
||||
# Check API keys
|
||||
if [ -z "${ANTHROPIC_API_KEY:-}" ]; then
|
||||
warn "ANTHROPIC_API_KEY not set — Opus teacher experiments will fail"
|
||||
fi
|
||||
if [ -z "${OPENAI_API_KEY:-}" ]; then
|
||||
warn "OPENAI_API_KEY not set — GPT-5.4 teacher experiments will fail"
|
||||
fi
|
||||
if [ -z "${GOOGLE_API_KEY:-}" ]; then
|
||||
warn "GOOGLE_API_KEY not set — Gemini teacher experiments will fail"
|
||||
fi
|
||||
|
||||
# Check Ollama
|
||||
if ! ollama list &>/dev/null; then
|
||||
fail "Ollama not running. Start it first."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Check student models
|
||||
for model in qwen3.5:2b qwen3.5:9b qwen3.5:27b; do
|
||||
if ! ollama list 2>/dev/null | grep -q "$model"; then
|
||||
warn "Model $model not found in Ollama. Pull with: ollama pull $model"
|
||||
fi
|
||||
done
|
||||
|
||||
# Check distillation init
|
||||
if [ ! -d "$HOME/.openjarvis/learning" ]; then
|
||||
log "Running jarvis learning init..."
|
||||
uv run jarvis learning init
|
||||
fi
|
||||
|
||||
ok "Preflight complete"
|
||||
}
|
||||
|
||||
# ── Run a single distillation session ────────────────────────────────────────
|
||||
run_session() {
|
||||
local config_file=$1
|
||||
local experiment_name
|
||||
experiment_name=$(basename "$(dirname "$config_file")")
|
||||
local config_name
|
||||
config_name=$(basename "${config_file%.toml}")
|
||||
local output_dir="${RESULTS_DIR}/${experiment_name}/${config_name}"
|
||||
|
||||
# Skip if already completed
|
||||
if [ -f "${output_dir}/session/session.json" ]; then
|
||||
ok "SKIP ${experiment_name}/${config_name} (already done)"
|
||||
return 0
|
||||
fi
|
||||
|
||||
log "──────────────────────────────────────────────────────"
|
||||
log "Experiment: ${experiment_name}/${config_name}"
|
||||
log "Config: ${config_file}"
|
||||
log "Output: ${output_dir}"
|
||||
log "──────────────────────────────────────────────────────"
|
||||
|
||||
mkdir -p "${output_dir}"
|
||||
|
||||
# Extract metadata from config
|
||||
local teacher_model
|
||||
teacher_model=$(grep 'teacher_model' "$config_file" | head -1 | sed 's/.*= *"\(.*\)"/\1/')
|
||||
local student_model
|
||||
student_model=$(grep 'default_model' "$config_file" | head -1 | sed 's/.*= *"\(.*\)"/\1/')
|
||||
local benchmark
|
||||
benchmark=$(grep '^benchmark ' "$config_file" | head -1 | sed 's/.*= *"\(.*\)"/\1/')
|
||||
local data_config
|
||||
data_config=$(grep 'data_config' "$config_file" | head -1 | sed 's/.*= *"\(.*\)"/\1/')
|
||||
local iterative
|
||||
iterative=$(grep 'iterative_sessions' "$config_file" | head -1 | sed 's/.*= *//')
|
||||
|
||||
log "Teacher: ${teacher_model}"
|
||||
log "Student: ${student_model}"
|
||||
log "Data: ${data_config:-C2}"
|
||||
log "Iter: ${iterative:-1}"
|
||||
|
||||
# ── Step 1: Seed traces based on data config ─────────────────────────
|
||||
# (In a full implementation, this would filter/prepare the TraceStore
|
||||
# based on C1/C2/C3. For now we use whatever traces exist.)
|
||||
|
||||
# ── Step 2: Run distillation session ─────────────────────────────────
|
||||
local n_sessions=${iterative:-1}
|
||||
local session_num=1
|
||||
local prev_session_id=""
|
||||
|
||||
while [ "$session_num" -le "$n_sessions" ]; do
|
||||
log "Session ${session_num}/${n_sessions}..."
|
||||
|
||||
local session_output="${output_dir}/session_${session_num}"
|
||||
mkdir -p "${session_output}"
|
||||
|
||||
# Run the distillation session via Python
|
||||
# (jarvis learning run doesn't support all config params yet,
|
||||
# so we call the orchestrator directly)
|
||||
uv run python << PYEOF > "${session_output}/run.log" 2>&1 || true
|
||||
import json, os, shutil, sys
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.engine.cloud import CloudEngine
|
||||
from openjarvis.evals.backends.jarvis_direct import JarvisDirectBackend
|
||||
from openjarvis.traces.store import TraceStore
|
||||
from openjarvis.learning.distillation.checkpoint.store import CheckpointStore
|
||||
from openjarvis.learning.distillation.models import AutonomyMode
|
||||
from openjarvis.learning.distillation.orchestrator import DistillationOrchestrator
|
||||
from openjarvis.learning.distillation.storage.session_store import SessionStore
|
||||
from openjarvis.learning.distillation.student_runner import (
|
||||
VLLMStudentRunner,
|
||||
build_benchmark_samples_from_traces,
|
||||
)
|
||||
from openjarvis.learning.distillation.triggers import OnDemandTrigger
|
||||
from openjarvis.learning.optimize.feedback.judge import TraceJudge
|
||||
|
||||
home = Path(os.environ.get("OPENJARVIS_HOME", str(Path.home() / ".openjarvis")))
|
||||
|
||||
# Read config params
|
||||
teacher_model = "${teacher_model}"
|
||||
student_model = "${student_model}"
|
||||
autonomy = "auto"
|
||||
max_cost = float("$(grep 'max_cost_per_session_usd' "$config_file" | head -1 | sed 's/.*= *//')")
|
||||
max_tools = int("$(grep 'max_tool_calls_per_diagnosis' "$config_file" | head -1 | sed 's/.*= *//')")
|
||||
|
||||
# Real student runner via vLLM
|
||||
vllm_host = os.environ.get("VLLM_HOST", "http://localhost:8001")
|
||||
student_runner = VLLMStudentRunner(
|
||||
host=vllm_host,
|
||||
model=student_model,
|
||||
)
|
||||
|
||||
# Real judge via cloud LLM
|
||||
cloud_engine = CloudEngine()
|
||||
judge_backend = JarvisDirectBackend(engine_key="cloud")
|
||||
judge = TraceJudge(backend=judge_backend, model="gpt-5-mini-2025-08-07")
|
||||
|
||||
# Build benchmark samples from existing traces
|
||||
trace_store = TraceStore(home / "traces.db")
|
||||
benchmark_samples = build_benchmark_samples_from_traces(trace_store, limit=50)
|
||||
|
||||
orch = DistillationOrchestrator(
|
||||
teacher_engine=cloud_engine,
|
||||
teacher_model=teacher_model,
|
||||
trace_store=trace_store,
|
||||
benchmark_samples=benchmark_samples,
|
||||
student_runner=student_runner,
|
||||
judge=judge,
|
||||
session_store=SessionStore(home / "learning" / "learning.db"),
|
||||
checkpoint_store=CheckpointStore(home),
|
||||
openjarvis_home=home,
|
||||
autonomy_mode=AutonomyMode.AUTO,
|
||||
scorer=None,
|
||||
min_traces=10,
|
||||
max_cost_usd=max_cost,
|
||||
max_tool_calls=max_tools,
|
||||
)
|
||||
session = orch.run(OnDemandTrigger())
|
||||
|
||||
# Save results
|
||||
result = {
|
||||
"session_id": session.id,
|
||||
"status": session.status.value,
|
||||
"cost_usd": session.teacher_cost_usd,
|
||||
"edits_total": len(session.edit_outcomes),
|
||||
"edits_applied": len([o for o in session.edit_outcomes if o.status == "applied"]),
|
||||
"edits_rejected": len([o for o in session.edit_outcomes if o.status == "rejected_by_gate"]),
|
||||
"error": session.error,
|
||||
}
|
||||
Path("${session_output}/result.json").write_text(json.dumps(result, indent=2))
|
||||
|
||||
# Copy session artifacts
|
||||
sd = home / "learning" / "sessions" / session.id
|
||||
if sd.exists():
|
||||
shutil.copytree(sd, Path("${session_output}/artifacts"), dirs_exist_ok=True)
|
||||
|
||||
print(json.dumps(result, indent=2))
|
||||
PYEOF
|
||||
|
||||
# Check result
|
||||
if [ -f "${session_output}/result.json" ]; then
|
||||
local status
|
||||
status=$(python3 -c "import json; print(json.load(open('${session_output}/result.json'))['status'])")
|
||||
local cost
|
||||
cost=$(python3 -c "import json; print(f\"\${json.load(open('${session_output}/result.json'))['cost_usd']:.4f}\")")
|
||||
local applied
|
||||
applied=$(python3 -c "import json; print(json.load(open('${session_output}/result.json'))['edits_applied'])")
|
||||
|
||||
if [ "$status" = "completed" ]; then
|
||||
ok "Session ${session_num}: status=${status}, cost=\$${cost}, applied=${applied}"
|
||||
else
|
||||
warn "Session ${session_num}: status=${status}, cost=\$${cost}"
|
||||
fi
|
||||
else
|
||||
fail "Session ${session_num}: no result.json (check ${session_output}/run.log)"
|
||||
fi
|
||||
|
||||
session_num=$((session_num + 1))
|
||||
done
|
||||
|
||||
ok "Done: ${experiment_name}/${config_name}"
|
||||
}
|
||||
|
||||
# ── Run experiment group ─────────────────────────────────────────────────────
|
||||
run_experiment() {
|
||||
local exp_dir=$1
|
||||
local exp_name
|
||||
exp_name=$(basename "$exp_dir")
|
||||
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
log "EXPERIMENT GROUP: ${exp_name}"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
|
||||
local count=0
|
||||
local total
|
||||
total=$(ls "${exp_dir}"/*.toml 2>/dev/null | wc -l)
|
||||
|
||||
for config in "${exp_dir}"/*.toml; do
|
||||
[ -f "$config" ] || continue
|
||||
|
||||
# Apply filter if specified
|
||||
if [ -n "${FILTER}" ] && ! echo "$config" | grep -q "${FILTER}"; then
|
||||
continue
|
||||
fi
|
||||
|
||||
count=$((count + 1))
|
||||
log "[${count}/${total}] $(basename "$config")"
|
||||
run_session "$config"
|
||||
done
|
||||
|
||||
ok "Experiment group ${exp_name}: ${count} configs processed"
|
||||
}
|
||||
|
||||
# ── Main ─────────────────────────────────────────────────────────────────────
|
||||
main() {
|
||||
check_prereqs
|
||||
|
||||
log "Starting distillation experiments"
|
||||
log "Experiment filter: ${EXPERIMENT}"
|
||||
log "Config filter: ${FILTER:-none}"
|
||||
|
||||
local start_time
|
||||
start_time=$(date +%s)
|
||||
|
||||
if [ "$EXPERIMENT" = "all" ]; then
|
||||
# Run in priority order
|
||||
for exp in exp1a-teacher exp1b-budget exp1c-student \
|
||||
exp2a-gate exp2b-autonomy \
|
||||
exp3a-iterative exp3b-transfer; do
|
||||
if [ -d "${CONFIGS_DIR}/${exp}" ]; then
|
||||
run_experiment "${CONFIGS_DIR}/${exp}"
|
||||
fi
|
||||
done
|
||||
elif [ -d "${CONFIGS_DIR}/${EXPERIMENT}" ]; then
|
||||
run_experiment "${CONFIGS_DIR}/${EXPERIMENT}"
|
||||
else
|
||||
fail "Unknown experiment: ${EXPERIMENT}"
|
||||
echo "Available: exp1a-teacher exp1b-budget exp1c-student exp2a-gate exp2b-autonomy exp3a-iterative exp3b-transfer"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
local end_time
|
||||
end_time=$(date +%s)
|
||||
local elapsed=$((end_time - start_time))
|
||||
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
ok "All experiments complete in ${elapsed}s"
|
||||
log "Results in: ${RESULTS_DIR}/"
|
||||
log "━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━"
|
||||
}
|
||||
|
||||
main "$@"
|
||||
@@ -0,0 +1,375 @@
|
||||
#!/usr/bin/env bash
|
||||
# =============================================================================
|
||||
# Track B: GEPA/DSPy Agent Optimization
|
||||
# NeurIPS 2026 — Agent Optimization Experiments
|
||||
#
|
||||
# Runs GEPA and DSPy BootstrapFewShot across:
|
||||
# Models: qwen-9b, qwen-27b, qwen-35b
|
||||
# Benchmarks: toolcall15, pinchbench, taubench
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/experiments/run_track_b_gepa_dspy.sh
|
||||
# bash scripts/experiments/run_track_b_gepa_dspy.sh --model qwen-9b --benchmark pinchbench
|
||||
# bash scripts/experiments/run_track_b_gepa_dspy.sh --optimizer gepa
|
||||
# bash scripts/experiments/run_track_b_gepa_dspy.sh --optimizer dspy
|
||||
#
|
||||
# Expected runtime: ~2-4 hours per GEPA run, ~1-2 hours per DSPy run
|
||||
# Total wall-clock: ~12-18 hours (parallelized across GPUs)
|
||||
# Estimated API cost: ~$90 GEPA + ~$90 DSPy = ~$180 total
|
||||
#
|
||||
# =============================================================================
|
||||
# vLLM Serving Commands (run these BEFORE this script on the GPU node)
|
||||
# =============================================================================
|
||||
#
|
||||
# GPU 0 — Qwen-9B (1x GPU):
|
||||
# CUDA_VISIBLE_DEVICES=0 vllm serve Qwen/Qwen2.5-7B-Instruct \
|
||||
# --model Qwen/Qwen2.5-7B-Instruct \
|
||||
# --served-model-name qwen-9b \
|
||||
# --port 8001 --host 0.0.0.0 \
|
||||
# --max-model-len 32768 --gpu-memory-utilization 0.9 &
|
||||
#
|
||||
# GPU 1 — Qwen-27B (1-2x GPU):
|
||||
# CUDA_VISIBLE_DEVICES=1,2 vllm serve Qwen/Qwen2.5-32B-Instruct \
|
||||
# --model Qwen/Qwen2.5-32B-Instruct \
|
||||
# --served-model-name qwen-27b \
|
||||
# --port 8002 --host 0.0.0.0 \
|
||||
# --tensor-parallel-size 2 \
|
||||
# --max-model-len 32768 --gpu-memory-utilization 0.9 &
|
||||
#
|
||||
# GPU 3 — Qwen-35B (1-2x GPU):
|
||||
# CUDA_VISIBLE_DEVICES=3,4 vllm serve Qwen/Qwen2.5-72B-Instruct \
|
||||
# --model Qwen/Qwen2.5-72B-Instruct \
|
||||
# --served-model-name qwen-35b \
|
||||
# --port 8003 --host 0.0.0.0 \
|
||||
# --tensor-parallel-size 2 \
|
||||
# --max-model-len 32768 --gpu-memory-utilization 0.9 &
|
||||
#
|
||||
# Wait for all servers to be healthy:
|
||||
# sleep 60 && curl -s http://localhost:8001/health && \
|
||||
# curl -s http://localhost:8002/health && \
|
||||
# curl -s http://localhost:8003/health
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
|
||||
RESULTS_BASE="${REPO_ROOT}/results/neurips-2026/agent-optimization"
|
||||
LOG_DIR="${REPO_ROOT}/results/neurips-2026/logs"
|
||||
TIMESTAMP="$(date +%Y%m%d_%H%M%S)"
|
||||
|
||||
ALL_MODELS=(qwen-9b qwen-27b qwen-35b)
|
||||
ALL_BENCHMARKS=(toolcall15 pinchbench taubench)
|
||||
ALL_OPTIMIZERS=(gepa dspy)
|
||||
|
||||
# Override defaults with CLI flags
|
||||
FILTER_MODEL=""
|
||||
FILTER_BENCHMARK=""
|
||||
FILTER_OPTIMIZER=""
|
||||
|
||||
# GEPA settings
|
||||
GEPA_TRIALS=20
|
||||
GEPA_MAX_SAMPLES=50
|
||||
GEPA_OPTIMIZER_MODEL="claude-sonnet-4-6"
|
||||
|
||||
# DSPy settings
|
||||
DSPY_OPTIMIZER="BootstrapFewShotWithRandomSearch"
|
||||
DSPY_TEACHER_LM="claude-sonnet-4-6"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Parse CLI flags
|
||||
# ---------------------------------------------------------------------------
|
||||
while [[ $# -gt 0 ]]; do
|
||||
case "$1" in
|
||||
--model)
|
||||
FILTER_MODEL="$2"; shift 2 ;;
|
||||
--benchmark)
|
||||
FILTER_BENCHMARK="$2"; shift 2 ;;
|
||||
--optimizer)
|
||||
FILTER_OPTIMIZER="$2"; shift 2 ;;
|
||||
--gepa-trials)
|
||||
GEPA_TRIALS="$2"; shift 2 ;;
|
||||
--gepa-max-samples)
|
||||
GEPA_MAX_SAMPLES="$2"; shift 2 ;;
|
||||
--dspy-optimizer)
|
||||
DSPY_OPTIMIZER="$2"; shift 2 ;;
|
||||
-h|--help)
|
||||
sed -n '2,30p' "$0" | grep '^#' | sed 's/^# \?//'
|
||||
exit 0 ;;
|
||||
*)
|
||||
echo "Unknown flag: $1"; exit 1 ;;
|
||||
esac
|
||||
done
|
||||
|
||||
# Apply filters
|
||||
if [[ -n "$FILTER_MODEL" ]]; then
|
||||
ALL_MODELS=("$FILTER_MODEL")
|
||||
fi
|
||||
if [[ -n "$FILTER_BENCHMARK" ]]; then
|
||||
ALL_BENCHMARKS=("$FILTER_BENCHMARK")
|
||||
fi
|
||||
if [[ -n "$FILTER_OPTIMIZER" ]]; then
|
||||
ALL_OPTIMIZERS=("$FILTER_OPTIMIZER")
|
||||
fi
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Logging helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
mkdir -p "$LOG_DIR"
|
||||
LOG_FILE="${LOG_DIR}/track_b_${TIMESTAMP}.log"
|
||||
|
||||
log() {
|
||||
local level="$1"; shift
|
||||
local msg="[$(date '+%Y-%m-%d %H:%M:%S')] [$level] $*"
|
||||
echo "$msg"
|
||||
echo "$msg" >> "$LOG_FILE"
|
||||
}
|
||||
|
||||
log_info() { log "INFO " "$@"; }
|
||||
log_ok() { log "OK " "$@"; }
|
||||
log_warn() { log "WARN " "$@"; }
|
||||
log_error() { log "ERROR" "$@"; }
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Environment setup
|
||||
# ---------------------------------------------------------------------------
|
||||
setup_env() {
|
||||
log_info "=== Track B: GEPA/DSPy Optimization ==="
|
||||
log_info "Repo: $REPO_ROOT"
|
||||
log_info "Log: $LOG_FILE"
|
||||
log_info "Models: ${ALL_MODELS[*]}"
|
||||
log_info "Benchmarks: ${ALL_BENCHMARKS[*]}"
|
||||
log_info "Optimizers: ${ALL_OPTIMIZERS[*]}"
|
||||
echo ""
|
||||
|
||||
# Check we are in the repo root
|
||||
if [[ ! -f "${REPO_ROOT}/pyproject.toml" ]]; then
|
||||
log_error "pyproject.toml not found — is REPO_ROOT set correctly? ($REPO_ROOT)"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Install/sync dependencies
|
||||
log_info "Running uv sync..."
|
||||
cd "$REPO_ROOT"
|
||||
uv sync --extra dev 2>&1 | tail -5
|
||||
log_ok "uv sync complete"
|
||||
|
||||
# Check required API keys
|
||||
if [[ -z "${ANTHROPIC_API_KEY:-}" ]]; then
|
||||
log_error "ANTHROPIC_API_KEY is not set. Required for the optimizer teacher model."
|
||||
log_error " export ANTHROPIC_API_KEY=sk-ant-..."
|
||||
exit 1
|
||||
fi
|
||||
log_ok "ANTHROPIC_API_KEY is set"
|
||||
|
||||
# Optional: OpenAI key (used if teacher_lm is an OpenAI model)
|
||||
if [[ -z "${OPENAI_API_KEY:-}" ]]; then
|
||||
log_warn "OPENAI_API_KEY not set (only needed if using OpenAI teacher models)"
|
||||
fi
|
||||
|
||||
echo ""
|
||||
log_info "Model-to-port mapping (vLLM must be pre-started on these ports):"
|
||||
log_info " qwen-9b -> http://localhost:8001"
|
||||
log_info " qwen-27b -> http://localhost:8002"
|
||||
log_info " qwen-35b -> http://localhost:8003"
|
||||
echo ""
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Model port lookup
|
||||
# ---------------------------------------------------------------------------
|
||||
model_port() {
|
||||
case "$1" in
|
||||
qwen-9b) echo 8001 ;;
|
||||
qwen-27b) echo 8002 ;;
|
||||
qwen-35b) echo 8003 ;;
|
||||
*)
|
||||
log_error "Unknown model: $1"
|
||||
exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health check: verify the vLLM server for a model is reachable
|
||||
# ---------------------------------------------------------------------------
|
||||
check_server_health() {
|
||||
local model="$1"
|
||||
local port
|
||||
port="$(model_port "$model")"
|
||||
local url="http://localhost:${port}/health"
|
||||
if curl -sf "$url" > /dev/null 2>&1; then
|
||||
log_ok "vLLM server for $model is healthy at port $port"
|
||||
return 0
|
||||
else
|
||||
log_error "vLLM server for $model NOT reachable at $url"
|
||||
log_error "Start it with the vLLM commands in the script header."
|
||||
return 1
|
||||
fi
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GEPA optimization for one (model, benchmark) pair
|
||||
# ---------------------------------------------------------------------------
|
||||
run_gepa() {
|
||||
local model="$1"
|
||||
local bench="$2"
|
||||
local port
|
||||
port="$(model_port "$model")"
|
||||
local out_dir="${RESULTS_BASE}/gepa/${model}/${bench}"
|
||||
|
||||
log_info "--- GEPA: $model × $bench ---"
|
||||
log_info " Output dir: $out_dir"
|
||||
log_info " Trials: $GEPA_TRIALS Max-samples: $GEPA_MAX_SAMPLES"
|
||||
mkdir -p "$out_dir"
|
||||
|
||||
# Record start time
|
||||
local t0
|
||||
t0="$(date +%s)"
|
||||
|
||||
OPENAI_API_BASE="http://localhost:${port}/v1" \
|
||||
uv run jarvis optimize run \
|
||||
--benchmark "$bench" \
|
||||
--model "$model" \
|
||||
--optimizer-model "$GEPA_OPTIMIZER_MODEL" \
|
||||
--trials "$GEPA_TRIALS" \
|
||||
--max-samples "$GEPA_MAX_SAMPLES" \
|
||||
--output-dir "$out_dir" \
|
||||
2>&1 | tee -a "$LOG_FILE"
|
||||
|
||||
local exit_code=${PIPESTATUS[0]}
|
||||
local t1
|
||||
t1="$(date +%s)"
|
||||
local elapsed=$(( t1 - t0 ))
|
||||
|
||||
if [[ $exit_code -eq 0 ]]; then
|
||||
log_ok "GEPA $model/$bench done in ${elapsed}s"
|
||||
else
|
||||
log_error "GEPA $model/$bench FAILED (exit $exit_code) after ${elapsed}s"
|
||||
return $exit_code
|
||||
fi
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DSPy optimization for one (model, benchmark) pair
|
||||
# ---------------------------------------------------------------------------
|
||||
run_dspy() {
|
||||
local model="$1"
|
||||
local bench="$2"
|
||||
local port
|
||||
port="$(model_port "$model")"
|
||||
local out_dir="${RESULTS_BASE}/dspy/${model}/${bench}"
|
||||
|
||||
log_info "--- DSPy: $model × $bench ---"
|
||||
log_info " Output dir: $out_dir"
|
||||
log_info " Teleprompter: $DSPY_OPTIMIZER Teacher: $DSPY_TEACHER_LM"
|
||||
mkdir -p "$out_dir"
|
||||
|
||||
local t0
|
||||
t0="$(date +%s)"
|
||||
|
||||
OPENAI_API_BASE="http://localhost:${port}/v1" \
|
||||
uv run python - <<PYEOF 2>&1 | tee -a "$LOG_FILE"
|
||||
import sys
|
||||
from openjarvis.learning.agents.dspy_optimizer import DSPyAgentOptimizer
|
||||
from openjarvis.core.config import DSPyOptimizerConfig
|
||||
from openjarvis.traces.store import TraceStore
|
||||
|
||||
store = TraceStore()
|
||||
config = DSPyOptimizerConfig(
|
||||
optimizer="${DSPY_OPTIMIZER}",
|
||||
teacher_lm="${DSPY_TEACHER_LM}",
|
||||
config_dir="${out_dir}",
|
||||
benchmark="${bench}",
|
||||
agent_filter="${model}",
|
||||
)
|
||||
result = DSPyAgentOptimizer(config).optimize(store)
|
||||
print(f"DSPy result for ${model}/${bench}: {result}")
|
||||
if result.get("status") not in ("ok", "success", "done"):
|
||||
sys.exit(1)
|
||||
PYEOF
|
||||
|
||||
local exit_code=${PIPESTATUS[0]}
|
||||
local t1
|
||||
t1="$(date +%s)"
|
||||
local elapsed=$(( t1 - t0 ))
|
||||
|
||||
if [[ $exit_code -eq 0 ]]; then
|
||||
log_ok "DSPy $model/$bench done in ${elapsed}s"
|
||||
else
|
||||
log_error "DSPy $model/$bench FAILED (exit $exit_code) after ${elapsed}s"
|
||||
return $exit_code
|
||||
fi
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Summary: print result file locations
|
||||
# ---------------------------------------------------------------------------
|
||||
print_summary() {
|
||||
echo ""
|
||||
log_info "=== Track B Complete ==="
|
||||
log_info "Results written to:"
|
||||
for opt in "${ALL_OPTIMIZERS[@]}"; do
|
||||
for model in "${ALL_MODELS[@]}"; do
|
||||
for bench in "${ALL_BENCHMARKS[@]}"; do
|
||||
local out_dir="${RESULTS_BASE}/${opt}/${model}/${bench}"
|
||||
if [[ -d "$out_dir" ]]; then
|
||||
log_ok " $opt/$model/$bench -> $out_dir"
|
||||
else
|
||||
log_warn " $opt/$model/$bench -> MISSING ($out_dir)"
|
||||
fi
|
||||
done
|
||||
done
|
||||
done
|
||||
log_info "Full log: $LOG_FILE"
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Main
|
||||
# ---------------------------------------------------------------------------
|
||||
main() {
|
||||
setup_env
|
||||
|
||||
local failed=0
|
||||
|
||||
# Pre-flight: check vLLM servers for all target models
|
||||
log_info "Checking vLLM server health..."
|
||||
for model in "${ALL_MODELS[@]}"; do
|
||||
check_server_health "$model" || failed=$(( failed + 1 ))
|
||||
done
|
||||
if [[ $failed -gt 0 ]]; then
|
||||
log_error "$failed vLLM server(s) not reachable. Start them first (see header)."
|
||||
exit 1
|
||||
fi
|
||||
echo ""
|
||||
|
||||
# Run all requested (optimizer, model, benchmark) combinations
|
||||
for opt in "${ALL_OPTIMIZERS[@]}"; do
|
||||
log_info "=========================================="
|
||||
log_info "Optimizer: $opt"
|
||||
log_info "=========================================="
|
||||
for model in "${ALL_MODELS[@]}"; do
|
||||
for bench in "${ALL_BENCHMARKS[@]}"; do
|
||||
case "$opt" in
|
||||
gepa) run_gepa "$model" "$bench" || failed=$(( failed + 1 )) ;;
|
||||
dspy) run_dspy "$model" "$bench" || failed=$(( failed + 1 )) ;;
|
||||
*) log_error "Unknown optimizer: $opt"; failed=$(( failed + 1 )) ;;
|
||||
esac
|
||||
echo ""
|
||||
done
|
||||
done
|
||||
done
|
||||
|
||||
print_summary
|
||||
|
||||
if [[ $failed -gt 0 ]]; then
|
||||
log_error "$failed run(s) failed. Check log for details: $LOG_FILE"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log_ok "All Track B runs completed successfully."
|
||||
}
|
||||
|
||||
main "$@"
|
||||
@@ -0,0 +1,778 @@
|
||||
#!/usr/bin/env bash
|
||||
# =============================================================================
|
||||
# Track D: LoRA / SFT Intelligence Optimization
|
||||
# NeurIPS 2026 — Intelligence Optimization Experiments
|
||||
#
|
||||
# Runs LoRA and SFT fine-tuning across:
|
||||
# LoRA models: Qwen-2B, Qwen-9B, Qwen-27B
|
||||
# SFT models: Qwen-2B, Qwen-9B
|
||||
#
|
||||
# After training, runs fast-benchmark eval on every checkpoint.
|
||||
#
|
||||
# Usage:
|
||||
# bash scripts/experiments/run_track_d_lora_sft.sh
|
||||
# bash scripts/experiments/run_track_d_lora_sft.sh --method lora
|
||||
# bash scripts/experiments/run_track_d_lora_sft.sh --method sft --model qwen-2b
|
||||
# bash scripts/experiments/run_track_d_lora_sft.sh --skip-eval
|
||||
#
|
||||
# Expected runtime per run:
|
||||
# Qwen-2B LoRA (~4-8 h, 1x H100, GPU 0)
|
||||
# Qwen-9B LoRA (~8-16 h, 1x H100, GPU 1)
|
||||
# Qwen-27B LoRA (~16-24 h, 2x H100, GPUs 2-3)
|
||||
# Qwen-2B SFT (~8-16 h, 1x H100, GPU 4)
|
||||
# Qwen-9B SFT (~16-24 h, 2x H100, GPUs 5-6)
|
||||
#
|
||||
# Total wall-clock (all in parallel): ~24 h on a 7-8x H100 node
|
||||
# Total GPU-hours: ~100-200 H100-hours
|
||||
#
|
||||
# =============================================================================
|
||||
# GPU Allocation — recommended for a single 8x H100 node
|
||||
# =============================================================================
|
||||
#
|
||||
# Run D1 (Qwen-2B LoRA) on GPU 0:
|
||||
# CUDA_VISIBLE_DEVICES=0 bash scripts/experiments/run_track_d_lora_sft.sh \
|
||||
# --method lora --model qwen-2b &
|
||||
#
|
||||
# Run D2 (Qwen-9B LoRA) on GPU 1:
|
||||
# CUDA_VISIBLE_DEVICES=1 bash scripts/experiments/run_track_d_lora_sft.sh \
|
||||
# --method lora --model qwen-9b &
|
||||
#
|
||||
# Run D3 (Qwen-27B LoRA) on GPUs 2-3:
|
||||
# CUDA_VISIBLE_DEVICES=2,3 bash scripts/experiments/run_track_d_lora_sft.sh \
|
||||
# --method lora --model qwen-27b &
|
||||
#
|
||||
# Run D4 (Qwen-2B SFT) on GPU 4:
|
||||
# CUDA_VISIBLE_DEVICES=4 bash scripts/experiments/run_track_d_lora_sft.sh \
|
||||
# --method sft --model qwen-2b &
|
||||
#
|
||||
# Run D5 (Qwen-9B SFT) on GPUs 5-6:
|
||||
# CUDA_VISIBLE_DEVICES=5,6 bash scripts/experiments/run_track_d_lora_sft.sh \
|
||||
# --method sft --model qwen-9b &
|
||||
#
|
||||
# =============================================================================
|
||||
# Training datasets (downloaded automatically if HF_TOKEN is set)
|
||||
# =============================================================================
|
||||
#
|
||||
# Primary agentic traces:
|
||||
# - neulab/agent-data-collection
|
||||
# - GAIR/AgentInstruct
|
||||
#
|
||||
# Supplementary reasoning:
|
||||
# - GeneralThought-430K-filtered
|
||||
# - GLM-4.7-flash SFT traces (168K + 57K)
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)"
|
||||
RESULTS_BASE="${REPO_ROOT}/results/neurips-2026/intelligence-optimization"
|
||||
LOG_DIR="${REPO_ROOT}/results/neurips-2026/logs"
|
||||
TIMESTAMP="$(date +%Y%m%d_%H%M%S)"
|
||||
|
||||
# Default: run all methods and models
|
||||
ALL_LORA_MODELS=(qwen-2b qwen-9b qwen-27b)
|
||||
ALL_SFT_MODELS=(qwen-2b qwen-9b)
|
||||
FAST_BENCHMARKS=(toolcall15 pinchbench taubench)
|
||||
|
||||
# CLI overrides
|
||||
FILTER_METHOD=""
|
||||
FILTER_MODEL=""
|
||||
SKIP_EVAL=false
|
||||
|
||||
# Training hyperparameters (override per model below)
|
||||
LORA_RANK=64
|
||||
LORA_ALPHA=128
|
||||
LORA_DROPOUT=0.05
|
||||
LORA_EPOCHS=3
|
||||
LORA_LR=2e-4
|
||||
LORA_BATCH_SIZE=8
|
||||
LORA_GRAD_ACCUM=4
|
||||
|
||||
SFT_EPOCHS=3
|
||||
SFT_LR=1e-5
|
||||
SFT_BATCH_SIZE=4
|
||||
SFT_GRAD_ACCUM=8
|
||||
|
||||
MAX_SEQ_LEN=8192
|
||||
|
||||
# Fast-eval settings
|
||||
EVAL_MAX_SAMPLES=20
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HuggingFace model IDs
|
||||
# ---------------------------------------------------------------------------
|
||||
model_hf_id() {
|
||||
case "$1" in
|
||||
qwen-2b) echo "Qwen/Qwen2.5-1.5B-Instruct" ;;
|
||||
qwen-9b) echo "Qwen/Qwen2.5-7B-Instruct" ;;
|
||||
qwen-27b) echo "Qwen/Qwen2.5-32B-Instruct" ;;
|
||||
*)
|
||||
echo "ERROR: unknown model $1" >&2
|
||||
exit 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Parse CLI flags
|
||||
# ---------------------------------------------------------------------------
|
||||
while [[ $# -gt 0 ]]; do
|
||||
case "$1" in
|
||||
--method)
|
||||
FILTER_METHOD="$2"; shift 2 ;;
|
||||
--model)
|
||||
FILTER_MODEL="$2"; shift 2 ;;
|
||||
--skip-eval)
|
||||
SKIP_EVAL=true; shift ;;
|
||||
--lora-rank)
|
||||
LORA_RANK="$2"; shift 2 ;;
|
||||
--lora-epochs)
|
||||
LORA_EPOCHS="$2"; shift 2 ;;
|
||||
--sft-epochs)
|
||||
SFT_EPOCHS="$2"; shift 2 ;;
|
||||
-h|--help)
|
||||
sed -n '2,35p' "$0" | grep '^#' | sed 's/^# \?//'
|
||||
exit 0 ;;
|
||||
*)
|
||||
echo "Unknown flag: $1"; exit 1 ;;
|
||||
esac
|
||||
done
|
||||
|
||||
# Apply model filter
|
||||
if [[ -n "$FILTER_MODEL" ]]; then
|
||||
ALL_LORA_MODELS=("$FILTER_MODEL")
|
||||
ALL_SFT_MODELS=("$FILTER_MODEL")
|
||||
fi
|
||||
|
||||
# Apply method filter
|
||||
RUN_LORA=true
|
||||
RUN_SFT=true
|
||||
if [[ "$FILTER_METHOD" == "lora" ]]; then
|
||||
RUN_SFT=false
|
||||
elif [[ "$FILTER_METHOD" == "sft" ]]; then
|
||||
RUN_LORA=false
|
||||
fi
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Logging helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
mkdir -p "$LOG_DIR"
|
||||
LOG_FILE="${LOG_DIR}/track_d_${TIMESTAMP}.log"
|
||||
|
||||
log() {
|
||||
local level="$1"; shift
|
||||
local msg="[$(date '+%Y-%m-%d %H:%M:%S')] [$level] $*"
|
||||
echo "$msg"
|
||||
echo "$msg" >> "$LOG_FILE"
|
||||
}
|
||||
|
||||
log_info() { log "INFO " "$@"; }
|
||||
log_ok() { log "OK " "$@"; }
|
||||
log_warn() { log "WARN " "$@"; }
|
||||
log_error() { log "ERROR" "$@"; }
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Environment setup
|
||||
# ---------------------------------------------------------------------------
|
||||
setup_env() {
|
||||
log_info "=== Track D: LoRA/SFT Training ==="
|
||||
log_info "Repo: $REPO_ROOT"
|
||||
log_info "Log: $LOG_FILE"
|
||||
log_info "Results: $RESULTS_BASE"
|
||||
echo ""
|
||||
|
||||
# Check we are in the repo root
|
||||
if [[ ! -f "${REPO_ROOT}/pyproject.toml" ]]; then
|
||||
log_error "pyproject.toml not found — is REPO_ROOT set correctly? ($REPO_ROOT)"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Install/sync dependencies
|
||||
log_info "Running uv sync..."
|
||||
cd "$REPO_ROOT"
|
||||
uv sync --extra dev 2>&1 | tail -5
|
||||
log_ok "uv sync complete"
|
||||
|
||||
# HuggingFace token — required to download gated Qwen models
|
||||
if [[ -z "${HF_TOKEN:-}" ]]; then
|
||||
log_warn "HF_TOKEN is not set."
|
||||
log_warn " If Qwen models are gated, set: export HF_TOKEN=hf_..."
|
||||
log_warn " Or pre-download them with: huggingface-cli download <model>"
|
||||
else
|
||||
log_ok "HF_TOKEN is set"
|
||||
# Log in so huggingface_hub uses the token
|
||||
uv run python -c "
|
||||
import huggingface_hub
|
||||
huggingface_hub.login(token='${HF_TOKEN}', add_to_git_credential=False)
|
||||
print('Logged in to HuggingFace Hub')
|
||||
" 2>&1 | tee -a "$LOG_FILE"
|
||||
fi
|
||||
|
||||
# Check for GPU
|
||||
if ! command -v nvidia-smi &>/dev/null; then
|
||||
log_warn "nvidia-smi not found — ensure CUDA is available for training."
|
||||
else
|
||||
log_info "GPU status:"
|
||||
nvidia-smi --query-gpu=index,name,memory.total,memory.free \
|
||||
--format=csv,noheader 2>&1 | while IFS= read -r line; do
|
||||
log_info " $line"
|
||||
done
|
||||
fi
|
||||
|
||||
# Check for trl / peft / transformers (training stack)
|
||||
uv run python -c "
|
||||
import importlib, sys
|
||||
missing = []
|
||||
for pkg in ['transformers', 'peft', 'trl', 'datasets', 'accelerate', 'bitsandbytes']:
|
||||
if importlib.util.find_spec(pkg) is None:
|
||||
missing.append(pkg)
|
||||
if missing:
|
||||
print('MISSING packages:', missing)
|
||||
sys.exit(1)
|
||||
else:
|
||||
print('Training stack OK: transformers, peft, trl, datasets, accelerate, bitsandbytes')
|
||||
" 2>&1 | tee -a "$LOG_FILE" || {
|
||||
log_error "Some training packages are missing. Install with:"
|
||||
log_error " uv add transformers peft trl datasets accelerate bitsandbytes"
|
||||
exit 1
|
||||
}
|
||||
|
||||
echo ""
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Download / verify training dataset
|
||||
# ---------------------------------------------------------------------------
|
||||
prepare_dataset() {
|
||||
local method="$1"
|
||||
local model="$2"
|
||||
local data_dir="${REPO_ROOT}/results/neurips-2026/training-data"
|
||||
mkdir -p "$data_dir"
|
||||
|
||||
log_info "Preparing training dataset for $method/$model..."
|
||||
|
||||
uv run python - <<PYEOF 2>&1 | tee -a "$LOG_FILE"
|
||||
from datasets import load_dataset
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
data_dir = Path("${data_dir}")
|
||||
output_path = data_dir / "${method}_${model}_train.jsonl"
|
||||
|
||||
if output_path.exists():
|
||||
lines = output_path.read_text().count('\n')
|
||||
print(f"Dataset already exists: {output_path} ({lines} examples)")
|
||||
else:
|
||||
print("Downloading neulab/agent-data-collection...")
|
||||
ds = load_dataset("neulab/agent-data-collection", split="train")
|
||||
print(f"Raw dataset size: {len(ds)}")
|
||||
|
||||
# Filter to reasonable-length examples for agentic fine-tuning
|
||||
examples = []
|
||||
for ex in ds:
|
||||
messages = ex.get("messages") or ex.get("conversations") or []
|
||||
if not messages:
|
||||
continue
|
||||
total_len = sum(len(str(m)) for m in messages)
|
||||
if 200 <= total_len <= 16000:
|
||||
examples.append({"messages": messages})
|
||||
|
||||
print(f"Filtered dataset size: {len(examples)}")
|
||||
|
||||
with output_path.open("w") as f:
|
||||
for ex in examples:
|
||||
f.write(json.dumps(ex) + "\n")
|
||||
print(f"Saved to {output_path}")
|
||||
PYEOF
|
||||
|
||||
echo "${data_dir}/${method}_${model}_train.jsonl"
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LoRA training for one model
|
||||
# ---------------------------------------------------------------------------
|
||||
run_lora() {
|
||||
local model="$1"
|
||||
local hf_id
|
||||
hf_id="$(model_hf_id "$model")"
|
||||
local out_dir="${RESULTS_BASE}/lora/${model}"
|
||||
local checkpoint_dir="${out_dir}/checkpoint"
|
||||
mkdir -p "$out_dir" "$checkpoint_dir"
|
||||
|
||||
log_info "=========================================="
|
||||
log_info "LoRA training: $model ($hf_id)"
|
||||
log_info " Output: $out_dir"
|
||||
log_info " LoRA rank=$LORA_RANK alpha=$LORA_ALPHA dropout=$LORA_DROPOUT"
|
||||
log_info " Epochs=$LORA_EPOCHS LR=$LORA_LR batch=$LORA_BATCH_SIZE grad_accum=$LORA_GRAD_ACCUM"
|
||||
log_info "=========================================="
|
||||
|
||||
local dataset_path
|
||||
dataset_path="$(prepare_dataset lora "$model")"
|
||||
|
||||
local t0
|
||||
t0="$(date +%s)"
|
||||
|
||||
# Number of GPUs visible
|
||||
local n_gpus
|
||||
n_gpus="$(python3 -c "import torch; print(torch.cuda.device_count())" 2>/dev/null || echo 1)"
|
||||
log_info "Training on $n_gpus GPU(s)"
|
||||
|
||||
if [[ "$n_gpus" -gt 1 ]]; then
|
||||
LAUNCHER="uv run torchrun --nproc_per_node=$n_gpus"
|
||||
else
|
||||
LAUNCHER="uv run python"
|
||||
fi
|
||||
|
||||
$LAUNCHER - <<PYEOF 2>&1 | tee -a "$LOG_FILE"
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from peft import LoraConfig, TaskType, get_peft_model
|
||||
from transformers import (
|
||||
AutoModelForCausalLM,
|
||||
AutoTokenizer,
|
||||
DataCollatorForSeq2Seq,
|
||||
Trainer,
|
||||
TrainingArguments,
|
||||
)
|
||||
|
||||
hf_id = "${hf_id}"
|
||||
out_dir = "${out_dir}"
|
||||
ckpt = "${checkpoint_dir}"
|
||||
data_path = "${dataset_path}"
|
||||
|
||||
# ---- Tokenizer ----
|
||||
print(f"Loading tokenizer: {hf_id}")
|
||||
tokenizer = AutoTokenizer.from_pretrained(hf_id, trust_remote_code=True)
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
# ---- Dataset ----
|
||||
print(f"Loading dataset: {data_path}")
|
||||
raw = load_dataset("json", data_files=data_path, split="train")
|
||||
raw = raw.train_test_split(test_size=0.02, seed=42)
|
||||
|
||||
def tokenize(example):
|
||||
messages = example["messages"]
|
||||
text = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=False,
|
||||
)
|
||||
enc = tokenizer(text, truncation=True, max_length=${MAX_SEQ_LEN})
|
||||
enc["labels"] = enc["input_ids"].copy()
|
||||
return enc
|
||||
|
||||
print("Tokenizing dataset...")
|
||||
tok_ds = raw.map(tokenize, remove_columns=raw["train"].column_names, num_proc=4)
|
||||
print(f"Train: {len(tok_ds['train'])} Eval: {len(tok_ds['test'])}")
|
||||
|
||||
# ---- Model ----
|
||||
print(f"Loading model: {hf_id}")
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
hf_id,
|
||||
torch_dtype=torch.bfloat16,
|
||||
trust_remote_code=True,
|
||||
device_map="auto",
|
||||
)
|
||||
|
||||
# ---- LoRA ----
|
||||
lora_config = LoraConfig(
|
||||
task_type=TaskType.CAUSAL_LM,
|
||||
r=${LORA_RANK},
|
||||
lora_alpha=${LORA_ALPHA},
|
||||
lora_dropout=${LORA_DROPOUT},
|
||||
target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
|
||||
"gate_proj", "up_proj", "down_proj"],
|
||||
bias="none",
|
||||
)
|
||||
model = get_peft_model(model, lora_config)
|
||||
model.print_trainable_parameters()
|
||||
|
||||
# ---- Training args ----
|
||||
steps_per_epoch = math.ceil(len(tok_ds["train"]) / (${LORA_BATCH_SIZE} * ${LORA_GRAD_ACCUM}))
|
||||
total_steps = steps_per_epoch * ${LORA_EPOCHS}
|
||||
save_steps = max(50, steps_per_epoch // 2)
|
||||
|
||||
args = TrainingArguments(
|
||||
output_dir=ckpt,
|
||||
num_train_epochs=${LORA_EPOCHS},
|
||||
per_device_train_batch_size=${LORA_BATCH_SIZE},
|
||||
gradient_accumulation_steps=${LORA_GRAD_ACCUM},
|
||||
learning_rate=${LORA_LR},
|
||||
warmup_ratio=0.05,
|
||||
lr_scheduler_type="cosine",
|
||||
logging_steps=10,
|
||||
save_steps=save_steps,
|
||||
eval_strategy="steps",
|
||||
eval_steps=save_steps,
|
||||
load_best_model_at_end=True,
|
||||
metric_for_best_model="eval_loss",
|
||||
bf16=True,
|
||||
gradient_checkpointing=True,
|
||||
dataloader_num_workers=4,
|
||||
report_to="none",
|
||||
save_total_limit=3,
|
||||
)
|
||||
|
||||
# ---- Trainer ----
|
||||
collator = DataCollatorForSeq2Seq(tokenizer, model=model, padding=True)
|
||||
trainer = Trainer(
|
||||
model=model,
|
||||
args=args,
|
||||
train_dataset=tok_ds["train"],
|
||||
eval_dataset=tok_ds["test"],
|
||||
data_collator=collator,
|
||||
)
|
||||
|
||||
print(f"Starting LoRA training: {total_steps} total steps")
|
||||
trainer.train()
|
||||
|
||||
# Save final adapter
|
||||
final_adapter = os.path.join(out_dir, "lora_adapter_final")
|
||||
model.save_pretrained(final_adapter)
|
||||
tokenizer.save_pretrained(final_adapter)
|
||||
print(f"Final LoRA adapter saved to {final_adapter}")
|
||||
|
||||
# Save training metadata
|
||||
meta = {
|
||||
"model": hf_id,
|
||||
"method": "lora",
|
||||
"lora_rank": ${LORA_RANK},
|
||||
"lora_alpha": ${LORA_ALPHA},
|
||||
"epochs": ${LORA_EPOCHS},
|
||||
"train_examples": len(tok_ds["train"]),
|
||||
"final_train_loss": trainer.state.log_history[-1].get("loss"),
|
||||
}
|
||||
with open(os.path.join(out_dir, "training_meta.json"), "w") as f:
|
||||
import json; json.dump(meta, f, indent=2)
|
||||
print("Training metadata saved.")
|
||||
PYEOF
|
||||
|
||||
local exit_code=${PIPESTATUS[0]}
|
||||
local t1
|
||||
t1="$(date +%s)"
|
||||
local elapsed=$(( t1 - t0 ))
|
||||
|
||||
if [[ $exit_code -eq 0 ]]; then
|
||||
log_ok "LoRA $model done in ${elapsed}s (~$(( elapsed / 3600 ))h $(( (elapsed % 3600) / 60 ))m)"
|
||||
return 0
|
||||
else
|
||||
log_error "LoRA $model FAILED (exit $exit_code) after ${elapsed}s"
|
||||
return $exit_code
|
||||
fi
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SFT (full fine-tuning) for one model
|
||||
# ---------------------------------------------------------------------------
|
||||
run_sft() {
|
||||
local model="$1"
|
||||
local hf_id
|
||||
hf_id="$(model_hf_id "$model")"
|
||||
local out_dir="${RESULTS_BASE}/sft/${model}"
|
||||
local checkpoint_dir="${out_dir}/checkpoint"
|
||||
mkdir -p "$out_dir" "$checkpoint_dir"
|
||||
|
||||
log_info "=========================================="
|
||||
log_info "SFT training: $model ($hf_id)"
|
||||
log_info " Output: $out_dir"
|
||||
log_info " Epochs=$SFT_EPOCHS LR=$SFT_LR batch=$SFT_BATCH_SIZE grad_accum=$SFT_GRAD_ACCUM"
|
||||
log_info "=========================================="
|
||||
|
||||
local dataset_path
|
||||
dataset_path="$(prepare_dataset sft "$model")"
|
||||
|
||||
local t0
|
||||
t0="$(date +%s)"
|
||||
|
||||
local n_gpus
|
||||
n_gpus="$(python3 -c "import torch; print(torch.cuda.device_count())" 2>/dev/null || echo 1)"
|
||||
log_info "Training on $n_gpus GPU(s)"
|
||||
|
||||
if [[ "$n_gpus" -gt 1 ]]; then
|
||||
LAUNCHER="uv run torchrun --nproc_per_node=$n_gpus"
|
||||
else
|
||||
LAUNCHER="uv run python"
|
||||
fi
|
||||
|
||||
$LAUNCHER - <<PYEOF 2>&1 | tee -a "$LOG_FILE"
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from transformers import (
|
||||
AutoModelForCausalLM,
|
||||
AutoTokenizer,
|
||||
DataCollatorForSeq2Seq,
|
||||
Trainer,
|
||||
TrainingArguments,
|
||||
)
|
||||
|
||||
hf_id = "${hf_id}"
|
||||
out_dir = "${out_dir}"
|
||||
ckpt = "${checkpoint_dir}"
|
||||
data_path = "${dataset_path}"
|
||||
|
||||
# ---- Tokenizer ----
|
||||
print(f"Loading tokenizer: {hf_id}")
|
||||
tokenizer = AutoTokenizer.from_pretrained(hf_id, trust_remote_code=True)
|
||||
if tokenizer.pad_token is None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
# ---- Dataset ----
|
||||
print(f"Loading dataset: {data_path}")
|
||||
raw = load_dataset("json", data_files=data_path, split="train")
|
||||
raw = raw.train_test_split(test_size=0.02, seed=42)
|
||||
|
||||
def tokenize(example):
|
||||
messages = example["messages"]
|
||||
text = tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=False,
|
||||
)
|
||||
enc = tokenizer(text, truncation=True, max_length=${MAX_SEQ_LEN})
|
||||
enc["labels"] = enc["input_ids"].copy()
|
||||
return enc
|
||||
|
||||
print("Tokenizing dataset...")
|
||||
tok_ds = raw.map(tokenize, remove_columns=raw["train"].column_names, num_proc=4)
|
||||
print(f"Train: {len(tok_ds['train'])} Eval: {len(tok_ds['test'])}")
|
||||
|
||||
# ---- Model (4-bit quantized to fit on fewer GPUs) ----
|
||||
print(f"Loading model: {hf_id}")
|
||||
from transformers import BitsAndBytesConfig
|
||||
bnb_config = BitsAndBytesConfig(
|
||||
load_in_4bit=True,
|
||||
bnb_4bit_compute_dtype=torch.bfloat16,
|
||||
bnb_4bit_quant_type="nf4",
|
||||
bnb_4bit_use_double_quant=True,
|
||||
)
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
hf_id,
|
||||
quantization_config=bnb_config,
|
||||
trust_remote_code=True,
|
||||
device_map="auto",
|
||||
)
|
||||
model.config.use_cache = False
|
||||
|
||||
# ---- Training args ----
|
||||
steps_per_epoch = math.ceil(len(tok_ds["train"]) / (${SFT_BATCH_SIZE} * ${SFT_GRAD_ACCUM}))
|
||||
total_steps = steps_per_epoch * ${SFT_EPOCHS}
|
||||
save_steps = max(50, steps_per_epoch // 2)
|
||||
|
||||
args = TrainingArguments(
|
||||
output_dir=ckpt,
|
||||
num_train_epochs=${SFT_EPOCHS},
|
||||
per_device_train_batch_size=${SFT_BATCH_SIZE},
|
||||
gradient_accumulation_steps=${SFT_GRAD_ACCUM},
|
||||
learning_rate=${SFT_LR},
|
||||
warmup_ratio=0.03,
|
||||
lr_scheduler_type="cosine",
|
||||
logging_steps=10,
|
||||
save_steps=save_steps,
|
||||
eval_strategy="steps",
|
||||
eval_steps=save_steps,
|
||||
load_best_model_at_end=True,
|
||||
metric_for_best_model="eval_loss",
|
||||
bf16=True,
|
||||
gradient_checkpointing=True,
|
||||
dataloader_num_workers=4,
|
||||
report_to="none",
|
||||
save_total_limit=3,
|
||||
)
|
||||
|
||||
# ---- Trainer ----
|
||||
collator = DataCollatorForSeq2Seq(tokenizer, model=model, padding=True)
|
||||
trainer = Trainer(
|
||||
model=model,
|
||||
args=args,
|
||||
train_dataset=tok_ds["train"],
|
||||
eval_dataset=tok_ds["test"],
|
||||
data_collator=collator,
|
||||
)
|
||||
|
||||
print(f"Starting SFT training: {total_steps} total steps")
|
||||
trainer.train()
|
||||
|
||||
# Save final model
|
||||
final_model = os.path.join(out_dir, "sft_model_final")
|
||||
model.save_pretrained(final_model)
|
||||
tokenizer.save_pretrained(final_model)
|
||||
print(f"Final SFT model saved to {final_model}")
|
||||
|
||||
# Save training metadata
|
||||
meta = {
|
||||
"model": hf_id,
|
||||
"method": "sft",
|
||||
"epochs": ${SFT_EPOCHS},
|
||||
"train_examples": len(tok_ds["train"]),
|
||||
"final_train_loss": trainer.state.log_history[-1].get("loss"),
|
||||
}
|
||||
with open(os.path.join(out_dir, "training_meta.json"), "w") as f:
|
||||
import json; json.dump(meta, f, indent=2)
|
||||
print("Training metadata saved.")
|
||||
PYEOF
|
||||
|
||||
local exit_code=${PIPESTATUS[0]}
|
||||
local t1
|
||||
t1="$(date +%s)"
|
||||
local elapsed=$(( t1 - t0 ))
|
||||
|
||||
if [[ $exit_code -eq 0 ]]; then
|
||||
log_ok "SFT $model done in ${elapsed}s (~$(( elapsed / 3600 ))h $(( (elapsed % 3600) / 60 ))m)"
|
||||
return 0
|
||||
else
|
||||
log_error "SFT $model FAILED (exit $exit_code) after ${elapsed}s"
|
||||
return $exit_code
|
||||
fi
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Post-training eval: run fast benchmarks on a checkpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
run_eval() {
|
||||
local method="$1" # lora or sft
|
||||
local model="$2"
|
||||
local out_dir="${RESULTS_BASE}/${method}/${model}"
|
||||
|
||||
if [[ "$SKIP_EVAL" == "true" ]]; then
|
||||
log_info "Skipping eval (--skip-eval)"
|
||||
return 0
|
||||
fi
|
||||
|
||||
# Locate the final checkpoint / adapter
|
||||
local checkpoint=""
|
||||
if [[ "$method" == "lora" ]]; then
|
||||
checkpoint="${out_dir}/lora_adapter_final"
|
||||
else
|
||||
checkpoint="${out_dir}/sft_model_final"
|
||||
fi
|
||||
|
||||
if [[ ! -d "$checkpoint" ]]; then
|
||||
log_warn "Checkpoint not found for $method/$model at $checkpoint — skipping eval"
|
||||
return 0
|
||||
fi
|
||||
|
||||
log_info "--- Post-training eval: $method/$model ---"
|
||||
log_info " Checkpoint: $checkpoint"
|
||||
|
||||
for bench in "${FAST_BENCHMARKS[@]}"; do
|
||||
local eval_out="${out_dir}/eval/${bench}"
|
||||
mkdir -p "$eval_out"
|
||||
log_info " Eval: $bench -> $eval_out"
|
||||
|
||||
uv run python - <<PYEOF 2>&1 | tee -a "$LOG_FILE"
|
||||
import subprocess, sys
|
||||
|
||||
cmd = [
|
||||
"uv", "run", "python", "-m", "openjarvis.evals", "run",
|
||||
"--model-path", "${checkpoint}",
|
||||
"--model-id", "${model}-${method}",
|
||||
"--benchmark", "${bench}",
|
||||
"--max-samples", "${EVAL_MAX_SAMPLES}",
|
||||
"--output", "${eval_out}",
|
||||
]
|
||||
print("Running:", " ".join(cmd))
|
||||
result = subprocess.run(cmd, capture_output=False)
|
||||
sys.exit(result.returncode)
|
||||
PYEOF
|
||||
|
||||
local eval_exit=${PIPESTATUS[0]}
|
||||
if [[ $eval_exit -eq 0 ]]; then
|
||||
log_ok " Eval $method/$model/$bench OK"
|
||||
else
|
||||
log_warn " Eval $method/$model/$bench returned exit $eval_exit (non-fatal)"
|
||||
fi
|
||||
done
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Summary
|
||||
# ---------------------------------------------------------------------------
|
||||
print_summary() {
|
||||
echo ""
|
||||
log_info "=== Track D Complete ==="
|
||||
log_info "Results written to:"
|
||||
|
||||
if [[ "$RUN_LORA" == "true" ]]; then
|
||||
for model in "${ALL_LORA_MODELS[@]}"; do
|
||||
local out="${RESULTS_BASE}/lora/${model}"
|
||||
if [[ -d "$out" ]]; then
|
||||
log_ok " lora/$model -> $out"
|
||||
else
|
||||
log_warn " lora/$model -> MISSING ($out)"
|
||||
fi
|
||||
done
|
||||
fi
|
||||
|
||||
if [[ "$RUN_SFT" == "true" ]]; then
|
||||
for model in "${ALL_SFT_MODELS[@]}"; do
|
||||
local out="${RESULTS_BASE}/sft/${model}"
|
||||
if [[ -d "$out" ]]; then
|
||||
log_ok " sft/$model -> $out"
|
||||
else
|
||||
log_warn " sft/$model -> MISSING ($out)"
|
||||
fi
|
||||
done
|
||||
fi
|
||||
|
||||
log_info "Full log: $LOG_FILE"
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Main
|
||||
# ---------------------------------------------------------------------------
|
||||
main() {
|
||||
setup_env
|
||||
|
||||
local failed=0
|
||||
|
||||
# ---- LoRA runs ----
|
||||
if [[ "$RUN_LORA" == "true" ]]; then
|
||||
log_info "=========================================="
|
||||
log_info "Starting LoRA training runs"
|
||||
log_info "Models: ${ALL_LORA_MODELS[*]}"
|
||||
log_info "=========================================="
|
||||
for model in "${ALL_LORA_MODELS[@]}"; do
|
||||
run_lora "$model" || { failed=$(( failed + 1 )); log_error "LoRA $model failed, continuing..."; }
|
||||
run_eval lora "$model"
|
||||
echo ""
|
||||
done
|
||||
fi
|
||||
|
||||
# ---- SFT runs ----
|
||||
if [[ "$RUN_SFT" == "true" ]]; then
|
||||
log_info "=========================================="
|
||||
log_info "Starting SFT training runs"
|
||||
log_info "Models: ${ALL_SFT_MODELS[*]}"
|
||||
log_info "=========================================="
|
||||
for model in "${ALL_SFT_MODELS[@]}"; do
|
||||
run_sft "$model" || { failed=$(( failed + 1 )); log_error "SFT $model failed, continuing..."; }
|
||||
run_eval sft "$model"
|
||||
echo ""
|
||||
done
|
||||
fi
|
||||
|
||||
print_summary
|
||||
|
||||
if [[ $failed -gt 0 ]]; then
|
||||
log_error "$failed training run(s) failed. Check log: $LOG_FILE"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
log_ok "All Track D runs completed successfully."
|
||||
}
|
||||
|
||||
main "$@"
|
||||
@@ -10,6 +10,10 @@ from __future__ import annotations
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from openjarvis.agents._stubs import AgentContext, AgentResult, ToolUsingAgent
|
||||
from openjarvis.agents.prompt_loader import (
|
||||
load_few_shot_exemplars,
|
||||
load_system_prompt_override,
|
||||
)
|
||||
from openjarvis.core.events import EventBus
|
||||
from openjarvis.core.registry import AgentRegistry
|
||||
from openjarvis.core.types import Message, Role, ToolCall, ToolResult
|
||||
@@ -217,9 +221,16 @@ class DeepResearchAgent(ToolUsingAgent):
|
||||
self._emit_turn_start(input)
|
||||
|
||||
# Build system prompt with current date/time injected
|
||||
messages = self._build_messages(
|
||||
input, context, system_prompt=_build_system_prompt()
|
||||
system_prompt = (
|
||||
load_system_prompt_override("deep_research") or _build_system_prompt()
|
||||
)
|
||||
messages = self._build_messages(input, context, system_prompt=system_prompt)
|
||||
|
||||
# Inject few-shot exemplars before the user input
|
||||
for ex in load_few_shot_exemplars("deep_research"):
|
||||
if ex.get("input") and ex.get("output"):
|
||||
messages.insert(-1, Message(role=Role.USER, content=ex["input"]))
|
||||
messages.insert(-1, Message(role=Role.ASSISTANT, content=ex["output"]))
|
||||
|
||||
# Prepare OpenAI-format tool definitions for native function calling
|
||||
tools_openai = [t.to_openai_function() for t in self._tools]
|
||||
|
||||
@@ -16,9 +16,14 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from openjarvis.agents._stubs import AgentContext, AgentResult, ToolUsingAgent
|
||||
from openjarvis.agents.prompt_loader import (
|
||||
load_few_shot_exemplars,
|
||||
load_system_prompt_override,
|
||||
)
|
||||
from openjarvis.core.events import EventBus
|
||||
from openjarvis.core.registry import AgentRegistry
|
||||
from openjarvis.core.types import Message, Role, ToolCall, ToolResult, _message_to_dict
|
||||
@@ -48,6 +53,15 @@ You are a Monitor Operative Agent designed for long-horizon tasks.
|
||||
2. STATE: Your previous findings and state are automatically restored
|
||||
3. MEMORY: Store important findings for future recall
|
||||
|
||||
## How to use tools
|
||||
|
||||
To call a tool, write on its own lines:
|
||||
|
||||
Action: <tool_name>
|
||||
Action Input: <json_arguments>
|
||||
|
||||
You will receive the result, then continue your response.
|
||||
|
||||
## Strategy
|
||||
- Memory extraction: {memory_extraction}
|
||||
- Observation compression: {observation_compression}
|
||||
@@ -169,14 +183,19 @@ class MonitorOperativeAgent(ToolUsingAgent):
|
||||
self._emit_turn_start(input)
|
||||
|
||||
# 1. Build system prompt with state context
|
||||
# Priority: constructor arg > file override > hardcoded default
|
||||
sys_parts: list[str] = []
|
||||
if self._system_prompt:
|
||||
sys_parts.append(self._system_prompt)
|
||||
else:
|
||||
tool_desc = self._build_tool_descriptions()
|
||||
prompt_template = (
|
||||
load_system_prompt_override("monitor_operative")
|
||||
or MONITOR_OPERATIVE_SYSTEM_PROMPT
|
||||
)
|
||||
try:
|
||||
sys_parts.append(
|
||||
MONITOR_OPERATIVE_SYSTEM_PROMPT.format(
|
||||
prompt_template.format(
|
||||
memory_extraction=self._memory_extraction,
|
||||
observation_compression=self._observation_compression,
|
||||
retrieval_strategy=self._retrieval_strategy,
|
||||
@@ -185,7 +204,7 @@ class MonitorOperativeAgent(ToolUsingAgent):
|
||||
),
|
||||
)
|
||||
except KeyError:
|
||||
sys_parts.append(MONITOR_OPERATIVE_SYSTEM_PROMPT)
|
||||
sys_parts.append(prompt_template)
|
||||
|
||||
# 2. State recall from memory backend
|
||||
previous_state = self._recall_state()
|
||||
@@ -205,6 +224,12 @@ class MonitorOperativeAgent(ToolUsingAgent):
|
||||
session_messages=session_messages,
|
||||
)
|
||||
|
||||
# 4b. Inject few-shot exemplars before the user input
|
||||
for ex in load_few_shot_exemplars("monitor_operative"):
|
||||
if ex.get("input") and ex.get("output"):
|
||||
messages.insert(-1, Message(role=Role.USER, content=ex["input"]))
|
||||
messages.insert(-1, Message(role=Role.ASSISTANT, content=ex["output"]))
|
||||
|
||||
# 5. Run function-calling tool loop
|
||||
openai_tools = self._executor.get_openai_tools() if self._tools else []
|
||||
all_tool_results: list[ToolResult] = []
|
||||
@@ -229,34 +254,59 @@ class MonitorOperativeAgent(ToolUsingAgent):
|
||||
for k in total_usage:
|
||||
total_usage[k] += usage.get(k, 0)
|
||||
content = result.get("content", "")
|
||||
# Strip think tags so they don't interfere with parsing
|
||||
content = self._strip_think_tags(content)
|
||||
raw_tool_calls = result.get("tool_calls", [])
|
||||
|
||||
# No tool calls -> check continuation, then final answer
|
||||
if not raw_tool_calls:
|
||||
# --- Native function-calling path ---
|
||||
if raw_tool_calls:
|
||||
tool_calls = [
|
||||
ToolCall(
|
||||
id=tc.get("id", f"call_{i}"),
|
||||
name=tc.get("name", ""),
|
||||
arguments=tc.get("arguments", "{}"),
|
||||
)
|
||||
for i, tc in enumerate(raw_tool_calls)
|
||||
]
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
else:
|
||||
# --- Text-based fallback ---
|
||||
tool_info = self._extract_tool_call(content)
|
||||
if tool_info:
|
||||
action, action_input = tool_info
|
||||
messages.append(Message(role=Role.ASSISTANT, content=content))
|
||||
tc = ToolCall(
|
||||
id=f"text_call_{turns}",
|
||||
name=action,
|
||||
arguments=action_input,
|
||||
)
|
||||
tool_result = self._executor.execute(tc)
|
||||
all_tool_results.append(tool_result)
|
||||
observation_content = self._compress_observation(
|
||||
tool_result.content
|
||||
)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content=f"Result: {observation_content}",
|
||||
)
|
||||
)
|
||||
self._extract_and_store(tc.name, tool_result.content)
|
||||
continue
|
||||
|
||||
# No tool calls at all -> check continuation, then final answer
|
||||
content = self._check_continuation(result, messages)
|
||||
break
|
||||
|
||||
# Build ToolCall objects from raw dicts
|
||||
tool_calls = [
|
||||
ToolCall(
|
||||
id=tc.get("id", f"call_{i}"),
|
||||
name=tc.get("name", ""),
|
||||
arguments=tc.get("arguments", "{}"),
|
||||
)
|
||||
for i, tc in enumerate(raw_tool_calls)
|
||||
]
|
||||
|
||||
# Append assistant message with tool calls
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content=content,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
)
|
||||
|
||||
# Execute each tool
|
||||
for tc in tool_calls:
|
||||
# Execute each native tool call
|
||||
tool_calls_to_exec = tool_calls
|
||||
for tc in tool_calls_to_exec:
|
||||
# Loop guard check
|
||||
if self._loop_guard:
|
||||
verdict = self._loop_guard.check_call(
|
||||
@@ -338,6 +388,89 @@ class MonitorOperativeAgent(ToolUsingAgent):
|
||||
metadata={**total_usage, "messages": msg_dicts},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Text-based tool call extraction (fallback for non-function-calling models)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_call(text: str) -> tuple[str, str] | None:
|
||||
"""Extract tool call from text output.
|
||||
|
||||
Supports three formats:
|
||||
1. Action: tool_name / Action Input: {"key": "value"}
|
||||
2. <tool_call>tool_name\\n$key=value</tool_call> (XML-style)
|
||||
3. <tool_name query="..."> or <tool_name>...</tool_name> (inline XML)
|
||||
"""
|
||||
# Format 1: Action / Action Input
|
||||
action_match = re.search(r"Action:\s*(.+)", text, re.IGNORECASE)
|
||||
input_match = re.search(
|
||||
r"Action Input:\s*(.+?)(?=\n\n|\Z)", text, re.DOTALL | re.IGNORECASE
|
||||
)
|
||||
if action_match:
|
||||
return (
|
||||
action_match.group(1).strip(),
|
||||
input_match.group(1).strip() if input_match else "{}",
|
||||
)
|
||||
|
||||
# Format 2: <tool_call>tool_name ... </tool_call>
|
||||
xml_match = re.search(
|
||||
r"<tool_call>\s*(\w+)\s*(.*?)</\w+>",
|
||||
text,
|
||||
re.DOTALL,
|
||||
)
|
||||
if xml_match:
|
||||
tool_name = xml_match.group(1).strip()
|
||||
raw_params = xml_match.group(2).strip()
|
||||
params: dict[str, Any] = {}
|
||||
for m in re.finditer(
|
||||
r"\$(\w+)=(.+?)(?=\$|\n<|</|$)", raw_params, re.DOTALL
|
||||
):
|
||||
params[m.group(1)] = m.group(2).strip().rstrip("</>\n")
|
||||
for m in re.finditer(r"<(\w+)>(.*?)</\1>", raw_params, re.DOTALL):
|
||||
key, val = m.group(1), m.group(2).strip()
|
||||
try:
|
||||
params[key] = int(val)
|
||||
except ValueError:
|
||||
params[key] = val
|
||||
if not params:
|
||||
for m in re.finditer(
|
||||
r"(\w+)\s*:\s*(.+?)(?=\n\w+\s*:|$)", raw_params, re.DOTALL
|
||||
):
|
||||
key, val = m.group(1), m.group(2).strip().strip("\"'")
|
||||
try:
|
||||
params[key] = int(val)
|
||||
except ValueError:
|
||||
params[key] = val
|
||||
if params:
|
||||
return (tool_name, json.dumps(params))
|
||||
return (tool_name, "{}")
|
||||
|
||||
# Format 3: <web_search query="..."> or <tool_name>args</tool_name>
|
||||
# Handles Qwen-style XML tool output like <web_search query="...">
|
||||
inline_match = re.search(
|
||||
r"<(\w+)\s+(.*?)/?>",
|
||||
text,
|
||||
re.DOTALL,
|
||||
)
|
||||
if inline_match:
|
||||
tool_name = inline_match.group(1).strip()
|
||||
# Skip common non-tool tags
|
||||
if tool_name.lower() in ("think", "br", "hr", "p", "div", "span", "b", "i"):
|
||||
return None
|
||||
attr_str = inline_match.group(2).strip()
|
||||
params = {}
|
||||
for m in re.finditer(r'(\w+)=["\']([^"\']*)["\']', attr_str):
|
||||
params[m.group(1)] = m.group(2)
|
||||
# Also handle unquoted: <web_search query=something>
|
||||
if not params:
|
||||
for m in re.finditer(r"(\w+)=(\S+)", attr_str):
|
||||
params[m.group(1)] = m.group(2).rstrip(">")
|
||||
if params:
|
||||
return (tool_name, json.dumps(params))
|
||||
return (tool_name, "{}")
|
||||
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Message building
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -12,6 +12,10 @@ import re
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from openjarvis.agents._stubs import AgentContext, AgentResult, ToolUsingAgent
|
||||
from openjarvis.agents.prompt_loader import (
|
||||
load_few_shot_exemplars,
|
||||
load_system_prompt_override,
|
||||
)
|
||||
from openjarvis.core.events import EventBus
|
||||
from openjarvis.core.registry import AgentRegistry
|
||||
from openjarvis.core.types import Message, Role, ToolCall, ToolResult
|
||||
@@ -220,7 +224,10 @@ class NativeOpenHandsAgent(ToolUsingAgent):
|
||||
self._emit_turn_start(input)
|
||||
|
||||
tool_descriptions = build_tool_descriptions(self._tools)
|
||||
system_prompt = OPENHANDS_SYSTEM_PROMPT.format(
|
||||
prompt_template = (
|
||||
load_system_prompt_override("native_openhands") or OPENHANDS_SYSTEM_PROMPT
|
||||
)
|
||||
system_prompt = prompt_template.format(
|
||||
tool_descriptions=tool_descriptions,
|
||||
)
|
||||
|
||||
@@ -276,6 +283,13 @@ class NativeOpenHandsAgent(ToolUsingAgent):
|
||||
)
|
||||
|
||||
messages = self._build_messages(input, context, system_prompt=system_prompt)
|
||||
|
||||
# Inject few-shot exemplars before the user input
|
||||
for ex in load_few_shot_exemplars("native_openhands"):
|
||||
if ex.get("input") and ex.get("output"):
|
||||
messages.insert(-1, Message(role=Role.USER, content=ex["input"]))
|
||||
messages.insert(-1, Message(role=Role.ASSISTANT, content=ex["output"]))
|
||||
|
||||
messages = self._truncate_if_needed(messages)
|
||||
|
||||
all_tool_results: list[ToolResult] = []
|
||||
|
||||
@@ -10,6 +10,10 @@ import re
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from openjarvis.agents._stubs import AgentContext, AgentResult, ToolUsingAgent
|
||||
from openjarvis.agents.prompt_loader import (
|
||||
load_few_shot_exemplars,
|
||||
load_system_prompt_override,
|
||||
)
|
||||
from openjarvis.core.events import EventBus
|
||||
from openjarvis.core.registry import AgentRegistry
|
||||
from openjarvis.core.types import Message, Role, ToolCall, ToolResult, _message_to_dict
|
||||
@@ -131,10 +135,8 @@ class NativeReActAgent(ToolUsingAgent):
|
||||
|
||||
# Build system prompt with rich tool descriptions
|
||||
tool_desc = build_tool_descriptions(self._tools)
|
||||
# Plan 2B I3: render the optimized few-shot examples as a section
|
||||
# that the model sees BEFORE the tool descriptions. When no
|
||||
# examples are present, this is an empty string and the rendered
|
||||
# prompt is unchanged.
|
||||
# Plan 2B I3: render optimized few-shot skill examples as a section
|
||||
# before the tool descriptions. Empty string when not present.
|
||||
if self._skill_few_shot_examples:
|
||||
skill_examples_block = (
|
||||
"## Skill Examples\n\n"
|
||||
@@ -143,13 +145,29 @@ class NativeReActAgent(ToolUsingAgent):
|
||||
)
|
||||
else:
|
||||
skill_examples_block = ""
|
||||
system_prompt = REACT_SYSTEM_PROMPT.format(
|
||||
tool_descriptions=tool_desc,
|
||||
skill_examples=skill_examples_block,
|
||||
# Respect $OPENJARVIS_HOME override for the base template (M2+ work).
|
||||
prompt_template = (
|
||||
load_system_prompt_override("native_react") or REACT_SYSTEM_PROMPT
|
||||
)
|
||||
# External overrides may not include the {skill_examples} slot.
|
||||
try:
|
||||
system_prompt = prompt_template.format(
|
||||
tool_descriptions=tool_desc,
|
||||
skill_examples=skill_examples_block,
|
||||
)
|
||||
except KeyError:
|
||||
system_prompt = prompt_template.format(tool_descriptions=tool_desc)
|
||||
if skill_examples_block:
|
||||
system_prompt = system_prompt + "\n\n" + skill_examples_block
|
||||
|
||||
messages = self._build_messages(input, context, system_prompt=system_prompt)
|
||||
|
||||
# Inject few-shot exemplars before the user input
|
||||
for ex in load_few_shot_exemplars("native_react"):
|
||||
if ex.get("input") and ex.get("output"):
|
||||
messages.insert(-1, Message(role=Role.USER, content=ex["input"]))
|
||||
messages.insert(-1, Message(role=Role.ASSISTANT, content=ex["output"]))
|
||||
|
||||
all_tool_results: list[ToolResult] = []
|
||||
turns = 0
|
||||
total_usage: dict[str, int] = {
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Load system prompt and few-shot overrides from $OPENJARVIS_HOME.
|
||||
|
||||
Distillation (M1) proposes edits that get written to disk by appliers.
|
||||
This module lets agents pick those overrides up at runtime:
|
||||
|
||||
- System prompts: ``$OPENJARVIS_HOME/agents/{name}/system_prompt.md``
|
||||
- Few-shot exemplars: ``$OPENJARVIS_HOME/agents/{name}/few_shot.json``
|
||||
|
||||
Override files are templates — they may contain ``{tool_descriptions}`` and
|
||||
other format placeholders that the agent fills in via ``.format()``, exactly
|
||||
like the hardcoded constants.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _openjarvis_home() -> Path:
|
||||
"""Resolve $OPENJARVIS_HOME, defaulting to ~/.openjarvis."""
|
||||
return Path(os.environ.get("OPENJARVIS_HOME", "~/.openjarvis")).expanduser()
|
||||
|
||||
|
||||
def load_system_prompt_override(agent_name: str) -> str | None:
|
||||
"""Return the override prompt for *agent_name*, or ``None``.
|
||||
|
||||
Looks for ``$OPENJARVIS_HOME/agents/<agent_name>/system_prompt.md``.
|
||||
``OPENJARVIS_HOME`` defaults to ``~/.openjarvis`` when unset.
|
||||
"""
|
||||
home = _openjarvis_home()
|
||||
prompt_path = home / "agents" / agent_name / "system_prompt.md"
|
||||
if not prompt_path.exists():
|
||||
return None
|
||||
try:
|
||||
content = prompt_path.read_text(encoding="utf-8")
|
||||
logger.info(
|
||||
"Loaded system prompt override for %s from %s", agent_name, prompt_path
|
||||
)
|
||||
return content
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to read system prompt override at %s", prompt_path, exc_info=True
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def load_few_shot_exemplars(
|
||||
agent_name: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return few-shot exemplars for *agent_name*, or empty list.
|
||||
|
||||
Looks for ``$OPENJARVIS_HOME/agents/<agent_name>/few_shot.json``.
|
||||
Expected format: ``[{"input": "Q", "output": "A"}, ...]``.
|
||||
"""
|
||||
home = _openjarvis_home()
|
||||
fs_path = home / "agents" / agent_name / "few_shot.json"
|
||||
if not fs_path.exists():
|
||||
return []
|
||||
try:
|
||||
data = json.loads(fs_path.read_text(encoding="utf-8"))
|
||||
if not isinstance(data, list):
|
||||
logger.warning("few_shot.json for %s is not a list", agent_name)
|
||||
return []
|
||||
logger.info(
|
||||
"Loaded %d few-shot exemplars for %s from %s",
|
||||
len(data),
|
||||
agent_name,
|
||||
fs_path,
|
||||
)
|
||||
return data
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to read few-shot exemplars at %s",
|
||||
fs_path,
|
||||
exc_info=True,
|
||||
)
|
||||
return []
|
||||
@@ -12,6 +12,10 @@ import re
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from openjarvis.agents._stubs import AgentContext, AgentResult, ToolUsingAgent
|
||||
from openjarvis.agents.prompt_loader import (
|
||||
load_few_shot_exemplars,
|
||||
load_system_prompt_override,
|
||||
)
|
||||
from openjarvis.agents.rlm_repl import RLMRepl
|
||||
from openjarvis.core.events import EventBus
|
||||
from openjarvis.core.registry import AgentRegistry
|
||||
@@ -154,13 +158,13 @@ class RLMAgent(ToolUsingAgent):
|
||||
if self._custom_system_prompt:
|
||||
system_prompt = self._custom_system_prompt
|
||||
else:
|
||||
prompt_template = load_system_prompt_override("rlm") or RLM_SYSTEM_PROMPT
|
||||
try:
|
||||
system_prompt = RLM_SYSTEM_PROMPT.format(
|
||||
system_prompt = prompt_template.format(
|
||||
tool_section=tool_section,
|
||||
)
|
||||
except KeyError:
|
||||
# Custom system_prompt override without {tool_section}
|
||||
system_prompt = RLM_SYSTEM_PROMPT
|
||||
system_prompt = prompt_template
|
||||
|
||||
# Create REPL with sub-LM callbacks
|
||||
repl = RLMRepl(
|
||||
@@ -181,6 +185,12 @@ class RLMAgent(ToolUsingAgent):
|
||||
system_prompt=system_prompt,
|
||||
)
|
||||
|
||||
# Inject few-shot exemplars before the user input
|
||||
for ex in load_few_shot_exemplars("rlm"):
|
||||
if ex.get("input") and ex.get("output"):
|
||||
messages.insert(-1, Message(role=Role.USER, content=ex["input"]))
|
||||
messages.insert(-1, Message(role=Role.ASSISTANT, content=ex["output"]))
|
||||
|
||||
all_tool_results: list[ToolResult] = []
|
||||
turns = 0
|
||||
total_usage: dict[str, int] = {
|
||||
|
||||
@@ -38,6 +38,7 @@ from openjarvis.cli.telemetry_cmd import telemetry
|
||||
from openjarvis.cli.tool_cmd import tool
|
||||
from openjarvis.cli.vault_cmd import vault
|
||||
from openjarvis.cli.workflow_cmd import workflow
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
|
||||
@click.group(help="OpenJarvis — modular AI assistant backend")
|
||||
@@ -98,6 +99,7 @@ cli.add_command(connect, "connect")
|
||||
cli.add_command(digest, "digest")
|
||||
cli.add_command(deep_research_setup, "deep-research-setup")
|
||||
cli.add_command(deep_research_setup, "research")
|
||||
cli.add_command(learning_group, "learning")
|
||||
|
||||
# Gateway CLI commands (lazy import to avoid pulling starlette)
|
||||
try:
|
||||
|
||||
@@ -29,7 +29,10 @@ class _OpenAICompatibleEngine(InferenceEngine):
|
||||
_api_prefix: str = "/v1"
|
||||
|
||||
def __init__(self, host: str | None = None, *, timeout: float = 600.0) -> None:
|
||||
self._host = (host or self._default_host).rstrip("/")
|
||||
import os
|
||||
|
||||
env_key = f"{self.engine_id.upper()}_HOST"
|
||||
self._host = (host or os.environ.get(env_key) or self._default_host).rstrip("/")
|
||||
self._client = httpx.Client(base_url=self._host, timeout=timeout)
|
||||
|
||||
# -- InferenceEngine interface ------------------------------------------
|
||||
|
||||
@@ -349,10 +349,10 @@ def _build_dataset(benchmark: str, subset: str | None = None):
|
||||
return LiveResearchBenchDataset(path=subset)
|
||||
elif benchmark == "liveresearchbench":
|
||||
from openjarvis.evals.datasets.liveresearchbench import (
|
||||
LiveResearchBenchSFDataset,
|
||||
LiveResearchBenchDataset as LRBDataset,
|
||||
)
|
||||
|
||||
return LiveResearchBenchSFDataset()
|
||||
return LRBDataset()
|
||||
elif benchmark == "toolcall15":
|
||||
from openjarvis.evals.datasets.toolcall15 import ToolCall15Dataset
|
||||
|
||||
@@ -507,10 +507,10 @@ def _build_scorer(benchmark: str, judge_backend, judge_model: str):
|
||||
return LiveResearchBenchScorer(judge_backend, judge_model)
|
||||
elif benchmark == "liveresearchbench":
|
||||
from openjarvis.evals.scorers.liveresearchbench import (
|
||||
LiveResearchBenchSFScorer,
|
||||
LiveResearchBenchScorer as LRBScorer,
|
||||
)
|
||||
|
||||
return LiveResearchBenchSFScorer(judge_backend, judge_model)
|
||||
return LRBScorer(judge_backend, judge_model)
|
||||
elif benchmark == "toolcall15":
|
||||
from openjarvis.evals.scorers.toolcall15 import ToolCall15Scorer
|
||||
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# LiveCodeBench eval: Qwen3.5-9B (vLLM, 1 GPU)
|
||||
# Competitive programming from LeetCode, AtCoder, CodeForces
|
||||
|
||||
[meta]
|
||||
name = "livecodebench-qwen-9b"
|
||||
description = "LiveCodeBench on Qwen/Qwen3.5-9B (vLLM, 1 GPU)"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-9b/livecodebench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-9B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "livecodebench"
|
||||
backend = "jarvis-direct"
|
||||
max_samples = 20
|
||||
@@ -0,0 +1,31 @@
|
||||
# LiveResearchBench: Qwen3.5-27B-FP8 (vLLM, 1 GPU)
|
||||
# Salesforce's checklist-based deep research benchmark.
|
||||
|
||||
[meta]
|
||||
name = "liveresearchbench-qwen-27b"
|
||||
description = "LiveResearchBench (Salesforce) on Qwen/Qwen3.5-27B-FP8"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.6
|
||||
max_tokens = 16384
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-27b/liveresearchbench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-27B-FP8"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "liveresearchbench"
|
||||
backend = "jarvis-direct"
|
||||
max_samples = 50
|
||||
@@ -0,0 +1,31 @@
|
||||
# LiveResearchBench: Qwen3.5-2B (vLLM, 1 GPU)
|
||||
# Salesforce's checklist-based deep research benchmark.
|
||||
|
||||
[meta]
|
||||
name = "liveresearchbench-qwen-2b"
|
||||
description = "LiveResearchBench (Salesforce) on Qwen/Qwen3.5-2B"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.6
|
||||
max_tokens = 16384
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-2b/liveresearchbench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-2B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "liveresearchbench"
|
||||
backend = "jarvis-direct"
|
||||
max_samples = 50
|
||||
@@ -0,0 +1,31 @@
|
||||
# LiveResearchBench: Qwen3.5-9B (vLLM, 1 GPU)
|
||||
# Salesforce's checklist-based deep research benchmark.
|
||||
|
||||
[meta]
|
||||
name = "liveresearchbench-qwen-9b"
|
||||
description = "LiveResearchBench (Salesforce) on Qwen/Qwen3.5-9B"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.6
|
||||
max_tokens = 16384
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-9b/liveresearchbench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-9B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "liveresearchbench"
|
||||
backend = "jarvis-direct"
|
||||
max_samples = 50
|
||||
@@ -0,0 +1,35 @@
|
||||
# PinchBench eval: Qwen3.5-9B (vLLM, 1 GPU)
|
||||
# Agent: native_openhands — all PinchBench-required tools enabled
|
||||
|
||||
[meta]
|
||||
name = "pinchbench-qwen-9b"
|
||||
description = "PinchBench on Qwen/Qwen3.5-9B (vLLM, 1 GPU)"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.6
|
||||
max_tokens = 8192
|
||||
|
||||
[judge]
|
||||
model = "claude-opus-4-5"
|
||||
temperature = 0.0
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-9b/pinchbench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-9B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "pinchbench"
|
||||
backend = "jarvis-agent"
|
||||
agent = "native_openhands"
|
||||
tools = [
|
||||
"think", "file_read", "file_write", "web_search", "shell_exec",
|
||||
"code_interpreter", "browser_navigate", "image_generate",
|
||||
"calculator", "http_request", "pdf_extract",
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
# TauBench V2 eval: Qwen3.5-9B (vLLM, 1 GPU)
|
||||
# Multi-turn customer service benchmark — airline + retail splits
|
||||
|
||||
[meta]
|
||||
name = "taubench-qwen-9b"
|
||||
description = "TauBench V2 on Qwen/Qwen3.5-9B (vLLM, 1 GPU)"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.7
|
||||
max_tokens = 4096
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-9b/taubench/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-9B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "taubench"
|
||||
backend = "jarvis-direct"
|
||||
max_samples = 20
|
||||
split = "airline,retail"
|
||||
@@ -0,0 +1,29 @@
|
||||
# ToolCall-15 eval: Qwen3.5-9B (vLLM, 1 GPU)
|
||||
# Lightweight tool calling benchmark — 15 scenarios, 5 categories
|
||||
|
||||
[meta]
|
||||
name = "toolcall15-qwen-9b"
|
||||
description = "ToolCall-15 on Qwen/Qwen3.5-9B (temperature=0)"
|
||||
|
||||
[defaults]
|
||||
temperature = 0.0
|
||||
max_tokens = 4096
|
||||
|
||||
[judge]
|
||||
model = "gpt-5-mini-2025-08-07"
|
||||
temperature = 0.0
|
||||
engine = "cloud"
|
||||
|
||||
[run]
|
||||
max_workers = 1
|
||||
output_dir = "results/neurips-2026/baselines/qwen-9b/toolcall15/"
|
||||
seed = 42
|
||||
|
||||
[[models]]
|
||||
name = "Qwen/Qwen3.5-9B"
|
||||
engine = "vllm"
|
||||
num_gpus = 1
|
||||
|
||||
[[benchmarks]]
|
||||
name = "toolcall15"
|
||||
backend = "jarvis-direct"
|
||||
@@ -64,6 +64,7 @@ KNOWN_BENCHMARKS = {
|
||||
"taubench",
|
||||
"livecodebench",
|
||||
"liveresearch",
|
||||
"liveresearchbench",
|
||||
"toolcall15",
|
||||
}
|
||||
|
||||
|
||||
@@ -1,20 +1,23 @@
|
||||
"""LiveResearchBench (Salesforce) dataset provider.
|
||||
"""LiveResearchBench dataset provider — Salesforce's checklist-based benchmark.
|
||||
|
||||
80 expert-curated deep research tasks with per-question evaluation
|
||||
checklists across three domains: daily life, enterprise, and academia.
|
||||
543 checklist items total (grouped by question).
|
||||
Loads Salesforce/LiveResearchBench from HuggingFace. Each task has a research
|
||||
question and a set of checklist items used for fine-grained, coverage-based
|
||||
evaluation.
|
||||
|
||||
Note: This is the actual LiveResearchBench by Salesforce (arxiv 2510.14240).
|
||||
The existing ``liveresearch`` module points at DeepResearchBench
|
||||
(Ayanami0730/deep_research_bench) despite its misleading class name.
|
||||
|
||||
Reference: https://github.com/SalesforceAIResearch/LiveResearchBench
|
||||
HuggingFace: Salesforce/LiveResearchBench (gated — accept terms first)
|
||||
Paper: https://arxiv.org/abs/2510.14240
|
||||
Dataset: https://huggingface.co/datasets/Salesforce/LiveResearchBench
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Iterable, List, Optional
|
||||
|
||||
from openjarvis.evals.core.dataset import DatasetProvider
|
||||
@@ -22,34 +25,59 @@ from openjarvis.evals.core.types import EvalRecord
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
_HF_DATASET = "Salesforce/LiveResearchBench"
|
||||
HF_DATASET_ID = "Salesforce/LiveResearchBench"
|
||||
DEFAULT_HF_CONFIG = "question_with_checklist"
|
||||
DEFAULT_HF_SPLIT = "test"
|
||||
|
||||
|
||||
def _replace_date_placeholders(text: str) -> str:
|
||||
"""Replace dynamic date placeholders in queries."""
|
||||
now = datetime.now()
|
||||
text = text.replace("{{current_year}}", str(now.year))
|
||||
text = text.replace("{{last_year}}", str(now.year - 1))
|
||||
text = text.replace("{{current_date}}", now.strftime("%Y-%m-%d"))
|
||||
text = text.replace("{{date}}", now.strftime("%Y-%m-%d"))
|
||||
text = re.sub(r"\{current_year\}", str(now.year), text)
|
||||
text = re.sub(r"\{last_year\}", str(now.year - 1), text)
|
||||
return text
|
||||
def _parse_checklist(checklist: Any) -> List[str]:
|
||||
"""Parse the checklist field — may be a JSON string, list, or newline text."""
|
||||
if isinstance(checklist, list):
|
||||
return [str(item).strip() for item in checklist if str(item).strip()]
|
||||
if not isinstance(checklist, str) or not checklist.strip():
|
||||
return []
|
||||
try:
|
||||
parsed = json.loads(checklist)
|
||||
if isinstance(parsed, list):
|
||||
return [str(item).strip() for item in parsed if str(item).strip()]
|
||||
if isinstance(parsed, str):
|
||||
return [parsed.strip()] if parsed.strip() else []
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# Fall back: treat as newline-separated items
|
||||
return [line.strip() for line in checklist.splitlines() if line.strip()]
|
||||
|
||||
|
||||
class LiveResearchBenchSFDataset(DatasetProvider):
|
||||
"""Salesforce LiveResearchBench — 80 expert-curated research tasks.
|
||||
class LiveResearchBenchDataset(DatasetProvider):
|
||||
"""LiveResearchBench — Salesforce's expert-curated deep research benchmark.
|
||||
|
||||
The HuggingFace dataset has 543 rows (multiple checklist items per
|
||||
question). We group by ``qid`` to produce one EvalRecord per unique
|
||||
question, with all checklist items aggregated in metadata.
|
||||
Loads tasks from HuggingFace with per-task checklists used for
|
||||
coverage-based evaluation. Tasks span 7 domains (Science/Tech, Business,
|
||||
Health, Law/Governance, Society/Culture, Education, Media).
|
||||
"""
|
||||
|
||||
dataset_id = "liveresearchbench"
|
||||
dataset_name = "LiveResearchBench (Salesforce)"
|
||||
dataset_name = "LiveResearchBench"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._records: Optional[List[EvalRecord]] = None
|
||||
def __init__(
|
||||
self,
|
||||
hf_config: Optional[str] = None,
|
||||
hf_split: Optional[str] = None,
|
||||
) -> None:
|
||||
self._hf_config = hf_config or DEFAULT_HF_CONFIG
|
||||
self._hf_split = hf_split or DEFAULT_HF_SPLIT
|
||||
self._records: List[EvalRecord] = []
|
||||
|
||||
def verify_requirements(self) -> List[str]:
|
||||
issues: List[str] = []
|
||||
try:
|
||||
import datasets # noqa: F401
|
||||
except ImportError:
|
||||
issues.append(
|
||||
"The 'datasets' package is required for LiveResearchBench. "
|
||||
"Install with: pip install datasets"
|
||||
)
|
||||
return issues
|
||||
|
||||
def load(
|
||||
self,
|
||||
@@ -58,122 +86,95 @@ class LiveResearchBenchSFDataset(DatasetProvider):
|
||||
split: Optional[str] = None,
|
||||
seed: Optional[int] = None,
|
||||
) -> None:
|
||||
try:
|
||||
from datasets import load_dataset
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
"datasets package required. Install with: pip install datasets"
|
||||
)
|
||||
from datasets import load_dataset
|
||||
|
||||
import os
|
||||
|
||||
hf_token = os.environ.get("HF_TOKEN") or os.environ.get(
|
||||
"HUGGING_FACE_HUB_TOKEN"
|
||||
)
|
||||
|
||||
# Try question_with_checklist first (has evaluation criteria)
|
||||
hf_config = split or "question_with_checklist"
|
||||
hf_split = split or self._hf_split
|
||||
LOGGER.info(
|
||||
"Loading LiveResearchBench from HuggingFace (%s, config=%s)",
|
||||
_HF_DATASET,
|
||||
hf_config,
|
||||
"Loading %s (config=%s, split=%s) from HuggingFace ...",
|
||||
HF_DATASET_ID,
|
||||
self._hf_config,
|
||||
hf_split,
|
||||
)
|
||||
ds = load_dataset(HF_DATASET_ID, self._hf_config, split=hf_split)
|
||||
|
||||
try:
|
||||
ds = load_dataset(
|
||||
_HF_DATASET, hf_config, split="test", token=hf_token
|
||||
)
|
||||
except Exception as exc:
|
||||
raise RuntimeError(
|
||||
f"Failed to load {_HF_DATASET}. This is a gated dataset — "
|
||||
"visit https://huggingface.co/datasets/Salesforce/LiveResearchBench "
|
||||
"to accept the terms, then set HF_TOKEN in your environment. "
|
||||
f"Error: {exc}"
|
||||
) from exc
|
||||
|
||||
# Group rows by qid (multiple checklist items per question)
|
||||
questions: Dict[str, Dict[str, Any]] = {}
|
||||
checklists_by_qid: Dict[str, List[str]] = defaultdict(list)
|
||||
|
||||
# `question_with_checklist` has multiple rows per qid (one per
|
||||
# checklist_id). Group by qid and fold all checklist items into a
|
||||
# single record per question.
|
||||
grouped: Dict[str, Dict[str, Any]] = {}
|
||||
for row in ds:
|
||||
qid = str(row.get("qid", ""))
|
||||
qid = str(row.get("qid", "") or "").strip()
|
||||
if not qid:
|
||||
continue
|
||||
|
||||
if qid not in questions:
|
||||
question = row.get("question", "") or row.get(
|
||||
"question_no_placeholder", ""
|
||||
)
|
||||
questions[qid] = {
|
||||
question = row.get("question") or row.get("question_no_placeholder") or ""
|
||||
if qid not in grouped:
|
||||
grouped[qid] = {
|
||||
"qid": qid,
|
||||
"question": question,
|
||||
"category": row.get("category", ""),
|
||||
"category": row.get("category", "") or "",
|
||||
"checklist": [],
|
||||
}
|
||||
|
||||
checklist = row.get("checklist", "") or row.get(
|
||||
"checklist_no_placeholder", ""
|
||||
cl_items = _parse_checklist(
|
||||
row.get("checklist") or row.get("checklist_no_placeholder")
|
||||
)
|
||||
if checklist:
|
||||
checklists_by_qid[qid].append(checklist)
|
||||
grouped[qid]["checklist"].extend(cl_items)
|
||||
|
||||
# Build EvalRecords
|
||||
records: List[EvalRecord] = []
|
||||
for qid, info in questions.items():
|
||||
question = info["question"]
|
||||
if not question:
|
||||
continue
|
||||
|
||||
question = _replace_date_placeholders(question)
|
||||
|
||||
problem = (
|
||||
"You are a research assistant. Please conduct thorough "
|
||||
"research on the following question and write a "
|
||||
"comprehensive report with citations.\n\n"
|
||||
f"{question}"
|
||||
)
|
||||
|
||||
metadata: Dict[str, Any] = {
|
||||
"qid": qid,
|
||||
"original_question": question,
|
||||
"category": info.get("category", ""),
|
||||
"checklists": checklists_by_qid.get(qid, []),
|
||||
}
|
||||
|
||||
records.append(
|
||||
EvalRecord(
|
||||
record_id=f"lrb-{qid}",
|
||||
problem=problem,
|
||||
reference="",
|
||||
category="liveresearchbench",
|
||||
metadata=metadata,
|
||||
)
|
||||
records = list(grouped.values())
|
||||
if not records:
|
||||
raise RuntimeError(
|
||||
f"LiveResearchBench: no records found in {HF_DATASET_ID} "
|
||||
f"(config={self._hf_config}, split={hf_split})"
|
||||
)
|
||||
|
||||
if seed is not None:
|
||||
rng = random.Random(seed)
|
||||
rng.shuffle(records)
|
||||
|
||||
random.Random(seed).shuffle(records)
|
||||
if max_samples is not None:
|
||||
records = records[:max_samples]
|
||||
|
||||
self._records = records
|
||||
total_checklists = sum(
|
||||
len(r.metadata.get("checklists", [])) for r in records
|
||||
)
|
||||
self._records = []
|
||||
for rec in records:
|
||||
question = (rec.get("question") or "").strip()
|
||||
if not question:
|
||||
LOGGER.warning("Skipping %s: empty question", rec.get("qid"))
|
||||
continue
|
||||
|
||||
research_prompt = (
|
||||
"You are a deep research assistant. Conduct thorough research "
|
||||
"on the following task and produce a comprehensive, "
|
||||
"well-structured, citation-grounded report that addresses "
|
||||
"every aspect of the request.\n\n"
|
||||
f"## Research Task\n\n{question}"
|
||||
)
|
||||
|
||||
self._records.append(
|
||||
EvalRecord(
|
||||
record_id=f"liveresearchbench-{rec['qid']}",
|
||||
problem=research_prompt,
|
||||
reference="", # checklist-based; no single reference answer
|
||||
category="liveresearchbench",
|
||||
subject=rec.get("category") or "",
|
||||
metadata={
|
||||
"qid": rec["qid"],
|
||||
"question": question,
|
||||
"checklist": rec["checklist"],
|
||||
"hf_category": rec.get("category", ""),
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
LOGGER.info(
|
||||
"LiveResearchBench: loaded %d tasks (%d checklist items)",
|
||||
"LiveResearchBench: loaded %d tasks (avg %.1f checklist items/task)",
|
||||
len(self._records),
|
||||
total_checklists,
|
||||
(
|
||||
sum(len(r.metadata.get("checklist", [])) for r in self._records)
|
||||
/ max(1, len(self._records))
|
||||
),
|
||||
)
|
||||
|
||||
def iter_records(self) -> Iterable[EvalRecord]:
|
||||
if self._records is None:
|
||||
raise RuntimeError("Call .load() before iterating")
|
||||
return iter(self._records)
|
||||
|
||||
def size(self) -> int:
|
||||
if self._records is None:
|
||||
raise RuntimeError("Call .load() before size()")
|
||||
return len(self._records)
|
||||
|
||||
|
||||
__all__ = ["LiveResearchBenchSFDataset"]
|
||||
__all__ = ["LiveResearchBenchDataset"]
|
||||
|
||||
@@ -78,6 +78,15 @@ def _tau2_to_oj_messages(
|
||||
return oj_msgs
|
||||
|
||||
|
||||
def _strip_think_tags(text: str) -> str:
|
||||
"""Remove ``<think>...</think>`` blocks from model output."""
|
||||
import re
|
||||
|
||||
text = re.sub(r"<think>.*?</think>\s*", "", text, flags=re.DOTALL | re.IGNORECASE)
|
||||
text = re.sub(r"^.*?</think>\s*", "", text, flags=re.DOTALL | re.IGNORECASE)
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _oj_result_to_tau2_msg(result: dict):
|
||||
"""Convert OpenJarvis engine.generate() result to a tau2 AssistantMessage."""
|
||||
from tau2.data_model.message import AssistantMessage, ToolCall
|
||||
@@ -99,9 +108,12 @@ def _oj_result_to_tau2_msg(result: dict):
|
||||
for tc in raw_tool_calls
|
||||
]
|
||||
|
||||
content = result.get("content") or ""
|
||||
content = _strip_think_tags(content) if content else None
|
||||
|
||||
return AssistantMessage(
|
||||
role="assistant",
|
||||
content=result.get("content") or None,
|
||||
content=content or None,
|
||||
tool_calls=tool_calls,
|
||||
cost=result.get("cost_usd", 0.0),
|
||||
)
|
||||
@@ -183,11 +195,17 @@ class JarvisHalfDuplexAgent:
|
||||
"tool_choice": "auto",
|
||||
}
|
||||
# Disable thinking mode for local models (Qwen3.5 etc.)
|
||||
# to avoid <think> tags interfering with tool call parsing
|
||||
# to avoid <think> tags interfering with tool call parsing.
|
||||
# vLLM >=0.8 accepts chat_template_kwargs as a top-level field.
|
||||
# We also set it inside extra_body for compatibility with
|
||||
# OpenAI SDK-based clients that only pass extra_body through.
|
||||
if "qwen" in self._model.lower():
|
||||
gen_kwargs["chat_template_kwargs"] = {
|
||||
"enable_thinking": False,
|
||||
}
|
||||
gen_kwargs.setdefault("extra_body", {})["chat_template_kwargs"] = {
|
||||
"enable_thinking": False,
|
||||
}
|
||||
result = self._engine.generate(oj_messages, **gen_kwargs)
|
||||
|
||||
# Convert result to tau2 AssistantMessage
|
||||
|
||||
@@ -1,70 +1,164 @@
|
||||
"""LiveResearchBench (Salesforce) scorer — checklist-based evaluation.
|
||||
"""LiveResearchBench scorer — checklist-based LLM-as-judge scoring.
|
||||
|
||||
Evaluates research reports using per-question checklists for coverage,
|
||||
plus LLM-as-judge for presentation quality and citation adequacy.
|
||||
For each task, LiveResearchBench provides a list of checklist items that a
|
||||
good response must cover. The scorer asks an LLM judge to evaluate each
|
||||
checklist item against the response, producing a per-item pass/fail.
|
||||
Final score = fraction of passed items (coverage).
|
||||
|
||||
Reference: https://github.com/SalesforceAIResearch/LiveResearchBench
|
||||
Reference: https://arxiv.org/abs/2510.14240
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from openjarvis.evals.core.scorer import LLMJudgeScorer
|
||||
from openjarvis.evals.core.types import EvalRecord
|
||||
|
||||
LOGGER = logging.getLogger(__name__)
|
||||
|
||||
_COVERAGE_PROMPT = """\
|
||||
You are evaluating a research report against a checklist of required topics/points.
|
||||
# Tasks with coverage >= PASS_THRESHOLD are marked as correct
|
||||
PASS_THRESHOLD = 0.5
|
||||
|
||||
**Research Question:**
|
||||
|
||||
def _build_judge_prompt(
|
||||
*,
|
||||
question: str,
|
||||
answer: str,
|
||||
checklist: List[str],
|
||||
) -> str:
|
||||
"""Build a batched judge prompt that evaluates all checklist items at once."""
|
||||
bullets = "\n".join(f"{i}. {item}" for i, item in enumerate(checklist, 1))
|
||||
return f"""You are an expert evaluator scoring a deep-research report against a checklist.
|
||||
|
||||
## Research Question
|
||||
{question}
|
||||
|
||||
**Report:**
|
||||
{report}
|
||||
## Checklist Items
|
||||
{bullets}
|
||||
|
||||
**Checklist items to verify (each should be covered in the report):**
|
||||
{checklist_items}
|
||||
## Report to Evaluate
|
||||
{answer}
|
||||
|
||||
For each checklist item, determine if the report adequately covers it.
|
||||
Respond with a JSON array of objects, one per checklist item:
|
||||
[
|
||||
{{"item": "<checklist item text>", "covered": true/false, "evidence": "<brief quote or reason>"}},
|
||||
...
|
||||
]
|
||||
## Instructions
|
||||
For each checklist item above, decide whether the report adequately covers it.
|
||||
A checklist item is COVERED if the report includes the required information or
|
||||
analysis in a clear, accurate, and substantive way. A checklist item is NOT
|
||||
COVERED if the report omits it, addresses it only superficially, or contains
|
||||
incorrect information.
|
||||
|
||||
Then on the last line, provide the overall coverage score:
|
||||
coverage_score: <number of covered items>/<total items>"""
|
||||
Return your evaluation as JSON with this exact structure:
|
||||
```json
|
||||
{{
|
||||
"judgments": [
|
||||
{{"index": 1, "covered": true, "reason": "brief justification"}},
|
||||
{{"index": 2, "covered": false, "reason": "brief justification"}}
|
||||
]
|
||||
}}
|
||||
```
|
||||
|
||||
_QUALITY_PROMPT = """\
|
||||
You are evaluating the quality of a research report.
|
||||
|
||||
**Research Question:**
|
||||
{question}
|
||||
|
||||
**Report:**
|
||||
{report}
|
||||
|
||||
Rate the report on these dimensions (each 1-5):
|
||||
|
||||
1. **Presentation**: Is the report well-structured, readable, and professional?
|
||||
2. **Depth**: Does the report go beyond surface-level information?
|
||||
3. **Citation**: Does the report reference specific sources, data, or evidence?
|
||||
4. **Consistency**: Is the report internally consistent and free of contradictions?
|
||||
|
||||
Respond in this exact format:
|
||||
presentation: <1-5>
|
||||
depth: <1-5>
|
||||
citation: <1-5>
|
||||
consistency: <1-5>
|
||||
reasoning: <brief explanation>"""
|
||||
Return one judgment per checklist item, indexed in the same order as above.
|
||||
Be rigorous: `covered: true` only if the report genuinely satisfies the item."""
|
||||
|
||||
|
||||
class LiveResearchBenchSFScorer(LLMJudgeScorer):
|
||||
"""Checklist + quality scorer for Salesforce LiveResearchBench."""
|
||||
def _parse_judge_response(raw: str, num_items: int) -> List[Dict[str, Any]]:
|
||||
"""Extract per-item judgments from the judge's raw response.
|
||||
|
||||
Returns one dict per checklist item ({index, covered, reason}). Missing
|
||||
items default to ``covered=False``.
|
||||
"""
|
||||
if not raw or not raw.strip():
|
||||
return [
|
||||
{"index": i + 1, "covered": False, "reason": "empty judge response"}
|
||||
for i in range(num_items)
|
||||
]
|
||||
|
||||
judgments: Dict[int, Dict[str, Any]] = {}
|
||||
|
||||
# Collect JSON candidates: code blocks first, then balanced braces.
|
||||
candidates: List[str] = []
|
||||
for match in re.finditer(r"```(?:json)?\s*(.*?)\s*```", raw, re.DOTALL):
|
||||
candidates.append(match.group(1))
|
||||
|
||||
depth = 0
|
||||
current: List[str] = []
|
||||
for char in raw:
|
||||
if char == "{":
|
||||
if depth == 0:
|
||||
current = []
|
||||
depth += 1
|
||||
if depth > 0:
|
||||
current.append(char)
|
||||
if char == "}":
|
||||
depth -= 1
|
||||
if depth == 0 and current:
|
||||
candidates.append("".join(current))
|
||||
|
||||
for candidate in candidates:
|
||||
try:
|
||||
parsed = json.loads(candidate)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not isinstance(parsed, dict):
|
||||
continue
|
||||
items = parsed.get("judgments") or parsed.get("items") or []
|
||||
if not isinstance(items, list):
|
||||
continue
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
idx_int = int(item.get("index"))
|
||||
except (ValueError, TypeError):
|
||||
continue
|
||||
covered_raw = item.get("covered")
|
||||
if isinstance(covered_raw, bool):
|
||||
covered = covered_raw
|
||||
elif isinstance(covered_raw, str):
|
||||
covered = covered_raw.strip().lower() in {
|
||||
"true",
|
||||
"yes",
|
||||
"covered",
|
||||
"1",
|
||||
"y",
|
||||
}
|
||||
else:
|
||||
covered = False
|
||||
judgments[idx_int] = {
|
||||
"index": idx_int,
|
||||
"covered": covered,
|
||||
"reason": str(item.get("reason", "")),
|
||||
}
|
||||
if judgments:
|
||||
break
|
||||
|
||||
if not judgments:
|
||||
LOGGER.warning("Failed to parse judge response — marking all items uncovered")
|
||||
|
||||
# Fill gaps with covered=False so downstream code always sees N items.
|
||||
return [
|
||||
judgments.get(
|
||||
i + 1,
|
||||
{
|
||||
"index": i + 1,
|
||||
"covered": False,
|
||||
"reason": "missing from judge output",
|
||||
},
|
||||
)
|
||||
for i in range(num_items)
|
||||
]
|
||||
|
||||
|
||||
class LiveResearchBenchScorer(LLMJudgeScorer):
|
||||
"""Checklist-based LLM-as-judge scorer for LiveResearchBench.
|
||||
|
||||
For each sample, the judge evaluates each checklist item against the
|
||||
model's report. Score = fraction covered. Tasks with score >=
|
||||
``PASS_THRESHOLD`` (default 0.5) are marked correct.
|
||||
"""
|
||||
|
||||
scorer_id = "liveresearchbench"
|
||||
|
||||
@@ -73,94 +167,62 @@ class LiveResearchBenchSFScorer(LLMJudgeScorer):
|
||||
record: EvalRecord,
|
||||
model_answer: str,
|
||||
) -> Tuple[Optional[bool], Dict[str, Any]]:
|
||||
checklist: List[str] = list(record.metadata.get("checklist", []) or [])
|
||||
|
||||
if not model_answer or not model_answer.strip():
|
||||
return False, {"reason": "empty_response", "score": 0.0}
|
||||
return False, {
|
||||
"score": 0.0,
|
||||
"coverage": 0.0,
|
||||
"covered_count": 0,
|
||||
"checklist_size": len(checklist),
|
||||
"reason": "empty_response",
|
||||
}
|
||||
|
||||
question = record.metadata.get("original_question", record.problem)
|
||||
checklists = record.metadata.get("checklists", [])
|
||||
|
||||
meta: Dict[str, Any] = {}
|
||||
|
||||
# Phase 1: Checklist coverage (if available)
|
||||
coverage_score = 0.0
|
||||
if checklists:
|
||||
try:
|
||||
checklist_text = "\n".join(
|
||||
f"- {item}" for item in checklists
|
||||
)
|
||||
prompt = _COVERAGE_PROMPT.format(
|
||||
question=question,
|
||||
report=model_answer[:8000], # Truncate long reports
|
||||
checklist_items=checklist_text,
|
||||
)
|
||||
raw = self._ask_judge(
|
||||
prompt, temperature=1.0, max_tokens=4096
|
||||
)
|
||||
|
||||
# Parse coverage_score from last line
|
||||
match = re.search(
|
||||
r"coverage_score:\s*(\d+)\s*/\s*(\d+)", raw
|
||||
)
|
||||
if match:
|
||||
covered = int(match.group(1))
|
||||
total = int(match.group(2))
|
||||
coverage_score = covered / total if total > 0 else 0.0
|
||||
meta["coverage_covered"] = covered
|
||||
meta["coverage_total"] = total
|
||||
else:
|
||||
# Fallback: count "covered": true occurrences
|
||||
covered = raw.lower().count('"covered": true') + raw.lower().count('"covered":true')
|
||||
total = len(checklists)
|
||||
coverage_score = covered / total if total > 0 else 0.0
|
||||
meta["coverage_covered"] = covered
|
||||
meta["coverage_total"] = total
|
||||
|
||||
meta["coverage_score"] = coverage_score
|
||||
meta["coverage_raw"] = raw[:500]
|
||||
except Exception as exc:
|
||||
LOGGER.warning(
|
||||
"Coverage scoring failed for %s: %s",
|
||||
record.record_id,
|
||||
exc,
|
||||
)
|
||||
meta["coverage_error"] = str(exc)
|
||||
|
||||
# Phase 2: Quality dimensions
|
||||
quality_score = 0.0
|
||||
try:
|
||||
prompt = _QUALITY_PROMPT.format(
|
||||
question=question,
|
||||
report=model_answer[:8000],
|
||||
)
|
||||
raw = self._ask_judge(prompt, temperature=1.0, max_tokens=1024)
|
||||
|
||||
dims = {}
|
||||
for dim in ["presentation", "depth", "citation", "consistency"]:
|
||||
match = re.search(rf"{dim}:\s*(\d)", raw)
|
||||
if match:
|
||||
dims[dim] = int(match.group(1))
|
||||
|
||||
if dims:
|
||||
quality_score = sum(dims.values()) / (5 * len(dims))
|
||||
meta["quality_dims"] = dims
|
||||
meta["quality_score"] = quality_score
|
||||
meta["quality_raw"] = raw[:500]
|
||||
except Exception as exc:
|
||||
if not checklist:
|
||||
LOGGER.warning(
|
||||
"Quality scoring failed for %s: %s", record.record_id, exc
|
||||
"No checklist attached to %s — scoring cannot produce a coverage number",
|
||||
record.record_id,
|
||||
)
|
||||
meta["quality_error"] = str(exc)
|
||||
return None, {
|
||||
"score": 0.0,
|
||||
"coverage": 0.0,
|
||||
"covered_count": 0,
|
||||
"checklist_size": 0,
|
||||
"reason": "no_checklist",
|
||||
}
|
||||
|
||||
# Final score: weighted average of coverage and quality
|
||||
if checklists:
|
||||
final_score = 0.6 * coverage_score + 0.4 * quality_score
|
||||
else:
|
||||
final_score = quality_score
|
||||
question = record.metadata.get("question") or record.problem
|
||||
|
||||
meta["final_score"] = final_score
|
||||
is_correct = final_score >= 0.5
|
||||
prompt = _build_judge_prompt(
|
||||
question=question,
|
||||
answer=model_answer,
|
||||
checklist=checklist,
|
||||
)
|
||||
|
||||
return is_correct, meta
|
||||
try:
|
||||
raw = self._ask_judge(prompt, temperature=0.0, max_tokens=4096)
|
||||
except Exception as exc:
|
||||
LOGGER.error("LLM judge call failed for %s: %s", record.record_id, exc)
|
||||
return None, {
|
||||
"score": 0.0,
|
||||
"error": str(exc),
|
||||
"checklist_size": len(checklist),
|
||||
}
|
||||
|
||||
judgments = _parse_judge_response(raw, num_items=len(checklist))
|
||||
covered_count = sum(1 for j in judgments if j["covered"])
|
||||
coverage = covered_count / len(checklist)
|
||||
is_correct = coverage >= PASS_THRESHOLD
|
||||
|
||||
metadata: Dict[str, Any] = {
|
||||
"score": coverage,
|
||||
"coverage": coverage,
|
||||
"covered_count": covered_count,
|
||||
"checklist_size": len(checklist),
|
||||
"judgments": judgments,
|
||||
"raw_judge_output": raw,
|
||||
}
|
||||
return is_correct, metadata
|
||||
|
||||
|
||||
__all__ = ["LiveResearchBenchSFScorer"]
|
||||
__all__ = ["LiveResearchBenchScorer"]
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Harness distillation: frontier-driven learning subsystem."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Git-backed checkpoint store for config rollback."""
|
||||
@@ -0,0 +1,255 @@
|
||||
"""Git-backed checkpoint store for distillation config rollback.
|
||||
|
||||
A thin wrapper over a local git repository at ``<openjarvis_home>/.git``.
|
||||
The repo tracks ``config.toml``, ``agents/``, and ``tools/`` so that the
|
||||
diff between two commits captures the harness state at any point in time.
|
||||
The repo does NOT track ``learning/`` (sessions are append-only artifacts,
|
||||
not config state).
|
||||
|
||||
The wrapper shells out to ``git`` via ``subprocess`` rather than depending
|
||||
on a third-party library — keeps the dependency surface zero.
|
||||
|
||||
See spec §7.4 and §7.6.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.storage.paths import (
|
||||
ConfigurationError,
|
||||
_find_source_root,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Files and directories the checkpoint repo tracks.
|
||||
_TRACKED_PATHS = ("config.toml", "agents", "tools")
|
||||
|
||||
# Marker line in baseline commits so we can detect them on re-init.
|
||||
_BASELINE_COMMIT_MESSAGE = "learning: checkpoint baseline"
|
||||
|
||||
|
||||
class DirtyWorkingTreeError(RuntimeError):
|
||||
"""Raised when begin_stage is called on a working tree with uncommitted changes."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StageHandle:
|
||||
"""Snapshot of repo state captured at the start of a staging operation."""
|
||||
|
||||
edit_id: str
|
||||
pre_stage_sha: str
|
||||
|
||||
|
||||
class CheckpointStore:
|
||||
"""Thin git wrapper for the distillation checkpoint repo.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
root :
|
||||
The directory that *contains* the checkpoint repo (the repo's
|
||||
``.git`` lives at ``root / ".git"``). For production use this is
|
||||
``~/.openjarvis/`` (or ``$OPENJARVIS_HOME``); for tests it's a
|
||||
``tmp_path`` subdirectory.
|
||||
"""
|
||||
|
||||
def __init__(self, root: Path) -> None:
|
||||
self._root = Path(root).resolve()
|
||||
|
||||
@property
|
||||
def root(self) -> Path:
|
||||
return self._root
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# init
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def init(self) -> None:
|
||||
"""Initialize the checkpoint repo if it doesn't exist.
|
||||
|
||||
Refuses to initialize if ``self.root`` is inside the OpenJarvis source
|
||||
tree — this is the same defense-in-depth check as
|
||||
``resolve_distillation_root``: we never want a stray git repo writing
|
||||
config snapshots into the working copy.
|
||||
|
||||
Idempotent: if ``.git`` already exists and contains a baseline commit,
|
||||
does nothing.
|
||||
"""
|
||||
source_root = _find_source_root()
|
||||
if source_root is not None:
|
||||
try:
|
||||
self._root.relative_to(source_root)
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
raise ConfigurationError(
|
||||
f"CheckpointStore root ({self._root}) is inside the "
|
||||
f"OpenJarvis source tree ({source_root}). Refusing to "
|
||||
"initialize a checkpoint repo there."
|
||||
)
|
||||
|
||||
self._root.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if (self._root / ".git").exists() and self._has_baseline_commit():
|
||||
return
|
||||
|
||||
if not (self._root / ".git").exists():
|
||||
self._git("init", "-q")
|
||||
self._git("config", "user.email", "distillation@openjarvis.local")
|
||||
self._git("config", "user.name", "OpenJarvis Distillation")
|
||||
|
||||
# Stage whatever tracked paths currently exist (it's OK if some
|
||||
# don't yet — the user may not have agents or tools dirs at first
|
||||
# init time, in which case the baseline commit is empty).
|
||||
for rel in _TRACKED_PATHS:
|
||||
target = self._root / rel
|
||||
if target.exists():
|
||||
self._git("add", rel)
|
||||
|
||||
# Allow empty so init succeeds even on a brand-new openjarvis home.
|
||||
self._git(
|
||||
"commit",
|
||||
"--allow-empty",
|
||||
"-q",
|
||||
"-m",
|
||||
_BASELINE_COMMIT_MESSAGE,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# current_sha
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def current_sha(self) -> str:
|
||||
"""Return the abbreviated sha of HEAD."""
|
||||
return self._git("rev-parse", "--short", "HEAD")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Staging primitives
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def begin_stage(self, edit_id: str) -> StageHandle:
|
||||
"""Capture HEAD sha and assert the working tree is clean.
|
||||
|
||||
Raises
|
||||
------
|
||||
DirtyWorkingTreeError
|
||||
If there are uncommitted changes to tracked files. The
|
||||
orchestrator should never start a stage on a dirty tree —
|
||||
it indicates the user has manual edits in progress.
|
||||
"""
|
||||
if self._working_tree_dirty():
|
||||
raise DirtyWorkingTreeError(
|
||||
f"Cannot begin stage for {edit_id}: working tree has "
|
||||
"uncommitted changes. Commit or stash them first."
|
||||
)
|
||||
return StageHandle(
|
||||
edit_id=edit_id,
|
||||
pre_stage_sha=self.current_sha(),
|
||||
)
|
||||
|
||||
def commit_stage(
|
||||
self,
|
||||
handle: StageHandle,
|
||||
*,
|
||||
message: str,
|
||||
session_id: str,
|
||||
risk_tier: str,
|
||||
) -> str:
|
||||
"""Stage all tracked-path changes and create a commit.
|
||||
|
||||
The commit message has the form::
|
||||
|
||||
<message>
|
||||
|
||||
Edit-ID: <handle.edit_id>
|
||||
Session-ID: <session_id>
|
||||
Risk-Tier: <risk_tier>
|
||||
|
||||
Returns the new commit sha (abbreviated).
|
||||
"""
|
||||
for rel in _TRACKED_PATHS:
|
||||
target = self._root / rel
|
||||
if target.exists():
|
||||
self._git("add", rel)
|
||||
|
||||
full_message = (
|
||||
f"{message}\n"
|
||||
"\n"
|
||||
f"Edit-ID: {handle.edit_id}\n"
|
||||
f"Session-ID: {session_id}\n"
|
||||
f"Risk-Tier: {risk_tier}\n"
|
||||
)
|
||||
self._git("commit", "-q", "-m", full_message)
|
||||
return self.current_sha()
|
||||
|
||||
def discard_stage(self, handle: StageHandle) -> None:
|
||||
"""Restore the working tree to ``handle.pre_stage_sha``."""
|
||||
for rel in _TRACKED_PATHS:
|
||||
target = self._root / rel
|
||||
if target.exists():
|
||||
self._git("checkout", handle.pre_stage_sha, "--", rel)
|
||||
# Sanity check: HEAD must still equal the pre-stage sha.
|
||||
if self.current_sha() != handle.pre_stage_sha:
|
||||
raise RuntimeError(
|
||||
"discard_stage left HEAD at unexpected sha — something "
|
||||
"committed during the stage. Aborting for safety."
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Session-level rollback
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def revert_session(self, session_id: str) -> list[str]:
|
||||
"""Revert all commits tagged ``Session-ID: <session_id>``.
|
||||
|
||||
Reverts are applied in reverse chronological order (newest first),
|
||||
each producing a *new* commit so history is preserved (no rewriting).
|
||||
Returns the list of new revert commit shas.
|
||||
"""
|
||||
log = self._git(
|
||||
"log",
|
||||
"--format=%H",
|
||||
f"--grep=Session-ID: {session_id}",
|
||||
)
|
||||
if not log.strip():
|
||||
return []
|
||||
|
||||
# log is newest-first, which is the order we want for revert.
|
||||
target_shas = log.strip().splitlines()
|
||||
new_shas: list[str] = []
|
||||
for sha in target_shas:
|
||||
self._git("revert", "--no-edit", sha)
|
||||
new_shas.append(self.current_sha())
|
||||
return new_shas
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internals
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _git(self, *args: str) -> str:
|
||||
"""Run a git command in the repo and return stripped stdout."""
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=self._root,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
)
|
||||
return result.stdout.strip()
|
||||
|
||||
def _has_baseline_commit(self) -> bool:
|
||||
try:
|
||||
log = self._git("log", "--format=%s")
|
||||
except subprocess.CalledProcessError:
|
||||
return False
|
||||
return _BASELINE_COMMIT_MESSAGE in log
|
||||
|
||||
def _working_tree_dirty(self) -> bool:
|
||||
"""Return True if there are uncommitted changes to tracked files."""
|
||||
status = self._git("status", "--porcelain")
|
||||
return bool(status.strip())
|
||||
@@ -0,0 +1,132 @@
|
||||
"""``jarvis learning`` — distillation learning CLI subcommands.
|
||||
|
||||
See spec §12.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import click
|
||||
from rich.console import Console
|
||||
|
||||
console = Console()
|
||||
|
||||
|
||||
@click.group("learning")
|
||||
def learning_group() -> None:
|
||||
"""Frontier-driven harness learning (distillation)."""
|
||||
|
||||
|
||||
@learning_group.command("init")
|
||||
def learning_init() -> None:
|
||||
"""Initialize the distillation checkpoint repo and directory layout."""
|
||||
from openjarvis.learning.distillation.checkpoint.store import CheckpointStore
|
||||
from openjarvis.learning.distillation.storage.paths import (
|
||||
ensure_distillation_dirs,
|
||||
resolve_distillation_root,
|
||||
)
|
||||
|
||||
root = resolve_distillation_root()
|
||||
ensure_distillation_dirs()
|
||||
home = root.parent # ~/.openjarvis
|
||||
store = CheckpointStore(home)
|
||||
store.init()
|
||||
console.print(f"[green]Initialized distillation at {root}[/green]")
|
||||
|
||||
|
||||
@learning_group.command("run")
|
||||
@click.option(
|
||||
"--autonomy",
|
||||
type=click.Choice(["auto", "tiered", "manual"]),
|
||||
default="tiered",
|
||||
)
|
||||
def learning_run(autonomy: str) -> None:
|
||||
"""Run an on-demand learning session."""
|
||||
console.print("[yellow]On-demand session started.[/yellow]")
|
||||
console.print("[dim]Use 'jarvis learning history' to check results.[/dim]")
|
||||
console.print(f"[dim]Autonomy mode: {autonomy}[/dim]")
|
||||
# Full wiring deferred to integration — this registers the CLI surface.
|
||||
console.print("[dim]Full orchestration requires configured teacher engine.[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("history")
|
||||
@click.option("--limit", type=int, default=10, help="Max sessions to show.")
|
||||
def learning_history(limit: int) -> None:
|
||||
"""List past learning sessions."""
|
||||
console.print(f"[dim]Showing last {limit} sessions (requires learning.db).[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("show")
|
||||
@click.argument("session_id")
|
||||
def learning_show(session_id: str) -> None:
|
||||
"""Show details of a learning session."""
|
||||
console.print(f"[dim]Session: {session_id}[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("review")
|
||||
def learning_review() -> None:
|
||||
"""Review pending edits awaiting approval."""
|
||||
console.print("[dim]Pending review queue.[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("approve")
|
||||
@click.argument("edit_id")
|
||||
def learning_approve(edit_id: str) -> None:
|
||||
"""Approve a pending edit."""
|
||||
console.print(f"[dim]Approving edit: {edit_id}[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("reject")
|
||||
@click.argument("edit_id")
|
||||
@click.option("--reason", type=str, default="", help="Rejection reason.")
|
||||
def learning_reject(edit_id: str, reason: str) -> None:
|
||||
"""Reject a pending edit."""
|
||||
console.print(f"[dim]Rejecting edit: {edit_id}[/dim]")
|
||||
|
||||
|
||||
@learning_group.command("rollback")
|
||||
@click.argument("session_id", required=False)
|
||||
@click.option("--last", is_flag=True, help="Rollback the most recent session.")
|
||||
def learning_rollback(session_id: str | None, last: bool) -> None:
|
||||
"""Rollback a learning session's commits."""
|
||||
target = session_id or ("last session" if last else "none")
|
||||
console.print(f"[dim]Rolling back: {target}[/dim]")
|
||||
|
||||
|
||||
@learning_group.group("benchmark")
|
||||
def benchmark_group() -> None:
|
||||
"""Personal benchmark management."""
|
||||
|
||||
|
||||
@benchmark_group.command("refresh")
|
||||
def benchmark_refresh() -> None:
|
||||
"""Manually refresh the personal benchmark."""
|
||||
console.print("[dim]Refreshing personal benchmark.[/dim]")
|
||||
|
||||
|
||||
@benchmark_group.command("show")
|
||||
def benchmark_show() -> None:
|
||||
"""Show current benchmark statistics."""
|
||||
console.print("[dim]Benchmark stats.[/dim]")
|
||||
|
||||
|
||||
@learning_group.group("daemon")
|
||||
def daemon_group() -> None:
|
||||
"""Background learning daemon."""
|
||||
|
||||
|
||||
@daemon_group.command("start")
|
||||
def daemon_start() -> None:
|
||||
"""Start the learning daemon."""
|
||||
console.print("[dim]Starting daemon.[/dim]")
|
||||
|
||||
|
||||
@daemon_group.command("stop")
|
||||
def daemon_stop() -> None:
|
||||
"""Stop the learning daemon."""
|
||||
console.print("[dim]Stopping daemon.[/dim]")
|
||||
|
||||
|
||||
@daemon_group.command("status")
|
||||
def daemon_status() -> None:
|
||||
"""Check daemon status."""
|
||||
console.print("[dim]Daemon status.[/dim]")
|
||||
@@ -0,0 +1 @@
|
||||
"""Diagnose phase: teacher-driven failure analysis."""
|
||||
@@ -0,0 +1,318 @@
|
||||
"""DiagnosisRunner: orchestrates phase 1 of the distillation loop.
|
||||
|
||||
Builds diagnostic tools, runs the TeacherAgent, parses failure clusters
|
||||
from the teacher's output, and persists artifacts.
|
||||
|
||||
See spec §5.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.learning.distillation.diagnose.teacher_agent import (
|
||||
TeacherAgent,
|
||||
)
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
from openjarvis.learning.distillation.diagnose.types import ToolCallRecord
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SYSTEM_PROMPT = """\
|
||||
You are a meta-engineer analyzing the performance of a local AI assistant \
|
||||
called OpenJarvis. Your job is to diagnose why the local student model fails \
|
||||
on certain tasks and identify root cause patterns.
|
||||
|
||||
IMPORTANT — OUTPUT REQUIREMENT: You MUST end your response with a JSON array \
|
||||
of failure clusters inside a ```json code fence. This is required. Example:
|
||||
|
||||
```json
|
||||
[
|
||||
{{
|
||||
"id": "cluster-001",
|
||||
"description": "Short description of this failure pattern",
|
||||
"sample_trace_ids": ["trace-abc", "trace-def", "trace-ghi"],
|
||||
"student_failure_rate": 0.75,
|
||||
"teacher_success_rate": 0.95,
|
||||
"skill_gap": "Explanation of the skill gap between student and teacher"
|
||||
}}
|
||||
]
|
||||
```
|
||||
|
||||
You have access to diagnostic tools that let you:
|
||||
- Browse and search the student's trace history
|
||||
- Read the student's current configuration, prompts, and tools
|
||||
- Re-run the student on benchmark tasks
|
||||
- Run yourself on the same tasks for comparison
|
||||
- Compare student vs teacher outputs
|
||||
|
||||
Your analysis should:
|
||||
1. Identify 2-5 failure clusters — groups of related failures with shared root causes.
|
||||
2. For each cluster, actually run the student and yourself on at least 3 sample tasks \
|
||||
to populate student_failure_rate and teacher_success_rate with real data.
|
||||
3. Describe the skill gap for each cluster.
|
||||
|
||||
Each cluster object in the final JSON array MUST have these fields:
|
||||
- id (string)
|
||||
- description (string)
|
||||
- sample_trace_ids (list of strings)
|
||||
- student_failure_rate (float 0-1)
|
||||
- teacher_success_rate (float 0-1)
|
||||
- skill_gap (string)
|
||||
|
||||
Budget: max ~{max_turns} tool calls, max ${max_cost_usd:.2f} USD.
|
||||
|
||||
Remember: You MUST end your response with the ```json ... ``` block described above.
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiagnosisResult:
|
||||
"""The output of a diagnosis run."""
|
||||
|
||||
diagnosis_md: str
|
||||
clusters: list[FailureCluster] = field(default_factory=list)
|
||||
cost_usd: float = 0.0
|
||||
tool_call_records: list[ToolCallRecord] = field(default_factory=list)
|
||||
|
||||
|
||||
class DiagnosisRunner:
|
||||
"""Orchestrates phase 1 of the distillation loop.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
teacher_engine :
|
||||
The CloudEngine (or mock) for teacher inference.
|
||||
teacher_model :
|
||||
Frontier model id (e.g. "claude-opus-4-6").
|
||||
trace_store :
|
||||
TraceStore for reading student traces.
|
||||
benchmark_samples :
|
||||
List of PersonalBenchmarkSample objects.
|
||||
student_runner :
|
||||
Callable to re-execute the student on a task.
|
||||
judge :
|
||||
TraceJudge for comparing outputs.
|
||||
session_dir :
|
||||
Path where session artifacts are written.
|
||||
session_id :
|
||||
Current session id.
|
||||
config :
|
||||
Dict with config_path and openjarvis_home.
|
||||
max_turns :
|
||||
Max teacher tool calls (default 30).
|
||||
max_cost_usd :
|
||||
Max teacher API cost (default 5.0).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
teacher_engine: Any,
|
||||
teacher_model: str,
|
||||
trace_store: Any,
|
||||
benchmark_samples: list,
|
||||
student_runner: Any,
|
||||
judge: Any,
|
||||
session_dir: Path,
|
||||
session_id: str,
|
||||
config: dict[str, Any],
|
||||
max_turns: int = 30,
|
||||
max_cost_usd: float = 5.0,
|
||||
) -> None:
|
||||
self._teacher_engine = teacher_engine
|
||||
self._teacher_model = teacher_model
|
||||
self._trace_store = trace_store
|
||||
self._benchmark_samples = benchmark_samples
|
||||
self._student_runner = student_runner
|
||||
self._judge = judge
|
||||
self._session_dir = Path(session_dir)
|
||||
self._session_id = session_id
|
||||
self._config = config
|
||||
self._max_turns = max_turns
|
||||
self._max_cost_usd = max_cost_usd
|
||||
|
||||
def run(self) -> DiagnosisResult:
|
||||
"""Execute the diagnosis phase.
|
||||
|
||||
Returns
|
||||
-------
|
||||
DiagnosisResult
|
||||
Contains the diagnosis markdown, parsed clusters, cost, and
|
||||
tool call records.
|
||||
"""
|
||||
# Ensure session directory exists
|
||||
self._session_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Build diagnostic tools
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=self._trace_store,
|
||||
config=self._config,
|
||||
benchmark_samples=self._benchmark_samples,
|
||||
student_runner=self._student_runner,
|
||||
teacher_engine=self._teacher_engine,
|
||||
teacher_model=self._teacher_model,
|
||||
judge=self._judge,
|
||||
session_id=self._session_id,
|
||||
)
|
||||
|
||||
# Build system prompt with budget hints
|
||||
system_prompt = _SYSTEM_PROMPT.format(
|
||||
max_turns=self._max_turns,
|
||||
max_cost_usd=self._max_cost_usd,
|
||||
)
|
||||
|
||||
# Run the teacher
|
||||
agent = TeacherAgent(
|
||||
engine=self._teacher_engine,
|
||||
model=self._teacher_model,
|
||||
tools=tools,
|
||||
max_turns=self._max_turns,
|
||||
max_cost_usd=self._max_cost_usd,
|
||||
)
|
||||
agent_result = agent.run(
|
||||
"Analyze the student's recent trace history, identify failure patterns, "
|
||||
"and produce a structured diagnosis with failure clusters.",
|
||||
system_prompt=system_prompt,
|
||||
)
|
||||
|
||||
# Persist diagnosis.md
|
||||
diagnosis_path = self._session_dir / "diagnosis.md"
|
||||
diagnosis_path.write_text(agent_result.content, encoding="utf-8")
|
||||
|
||||
# Persist teacher traces JSONL
|
||||
traces_dir = self._session_dir / "teacher_traces"
|
||||
traces_dir.mkdir(parents=True, exist_ok=True)
|
||||
jsonl_path = traces_dir / "diagnose.jsonl"
|
||||
with jsonl_path.open("w", encoding="utf-8") as f:
|
||||
for record in agent_result.tool_call_records:
|
||||
f.write(json.dumps(record.to_jsonl_dict()) + "\n")
|
||||
|
||||
# Parse failure clusters from the diagnosis content
|
||||
clusters = _parse_clusters(agent_result.content)
|
||||
|
||||
# Fallback: if no clusters parsed, ask the teacher to emit only the JSON
|
||||
if not clusters:
|
||||
logger.warning(
|
||||
"No clusters in primary diagnosis (%d chars, %d tool calls). "
|
||||
"Attempting fallback extraction.",
|
||||
len(agent_result.content),
|
||||
len(agent_result.tool_call_records),
|
||||
)
|
||||
clusters = self._fallback_extract_clusters(agent_result.content)
|
||||
|
||||
return DiagnosisResult(
|
||||
diagnosis_md=agent_result.content,
|
||||
clusters=clusters,
|
||||
cost_usd=agent_result.total_cost_usd,
|
||||
tool_call_records=agent_result.tool_call_records,
|
||||
)
|
||||
|
||||
def _fallback_extract_clusters(self, diagnosis: str) -> list[FailureCluster]:
|
||||
"""One-shot fallback: ask the teacher to emit only the JSON array.
|
||||
|
||||
Makes a single no-tools call with max 1 turn. If the response still
|
||||
does not contain valid clusters, returns an empty list.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
diagnosis :
|
||||
The full diagnosis text from the primary teacher run.
|
||||
"""
|
||||
fallback_prompt = (
|
||||
"Your diagnosis did not include the required JSON cluster array. "
|
||||
"Here is your diagnosis:\n\n"
|
||||
f"{diagnosis[:3000]}\n\n"
|
||||
"Now output ONLY a raw JSON array of failure clusters. "
|
||||
"Each object must have: id, description, sample_trace_ids, "
|
||||
"student_failure_rate (float 0-1), teacher_success_rate (float 0-1), "
|
||||
"skill_gap. Output ONLY the JSON array — no markdown, no code "
|
||||
"fences, no explanation, no other text."
|
||||
)
|
||||
|
||||
fallback_agent = TeacherAgent(
|
||||
engine=self._teacher_engine,
|
||||
model=self._teacher_model,
|
||||
tools=[], # no tools — single generation call
|
||||
max_turns=1,
|
||||
max_cost_usd=self._max_cost_usd,
|
||||
)
|
||||
try:
|
||||
fallback_result = fallback_agent.run(fallback_prompt)
|
||||
except Exception:
|
||||
logger.exception("Fallback extraction TeacherAgent call failed")
|
||||
return []
|
||||
|
||||
clusters = _parse_clusters(fallback_result.content)
|
||||
if not clusters:
|
||||
logger.warning("Fallback extraction also produced no clusters")
|
||||
return clusters
|
||||
|
||||
|
||||
def _parse_clusters(content: str) -> list[FailureCluster]:
|
||||
"""Extract failure clusters from teacher diagnosis output.
|
||||
|
||||
Looks for a JSON array inside a ```json code fence. Falls back to
|
||||
searching for any JSON array in the content. Returns an empty list
|
||||
if no valid clusters are found.
|
||||
"""
|
||||
# Try to find JSON in a code fence first
|
||||
fence_match = re.search(r"```json\s*\n(.*?)\n```", content, re.DOTALL)
|
||||
if fence_match:
|
||||
try:
|
||||
return _parse_cluster_list(fence_match.group(1))
|
||||
except Exception as e:
|
||||
logger.warning("Failed to parse clusters from JSON code fence: %s", e)
|
||||
|
||||
# Fallback: try to find any JSON array in the content
|
||||
# Use greedy match to capture the full array (not just the first [...])
|
||||
for match in re.finditer(r"\[[\s\S]*\]", content):
|
||||
try:
|
||||
return _parse_cluster_list(match.group(0))
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
logger.warning("No failure clusters found in diagnosis output")
|
||||
return []
|
||||
|
||||
|
||||
def _parse_cluster_list(json_str: str) -> list[FailureCluster]:
|
||||
"""Parse a JSON string into a list of FailureCluster."""
|
||||
data = json.loads(json_str)
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
clusters = []
|
||||
for item in data:
|
||||
try:
|
||||
# Clamp rates to 0-1 range (teacher sometimes outputs percentages)
|
||||
failure_rate = float(item.get("student_failure_rate", 0))
|
||||
success_rate = float(item.get("teacher_success_rate", 0))
|
||||
if failure_rate > 1.0:
|
||||
failure_rate = failure_rate / 100.0
|
||||
if success_rate > 1.0:
|
||||
success_rate = success_rate / 100.0
|
||||
failure_rate = max(0.0, min(1.0, failure_rate))
|
||||
success_rate = max(0.0, min(1.0, success_rate))
|
||||
|
||||
clusters.append(
|
||||
FailureCluster(
|
||||
id=str(item.get("id", f"cluster-{len(clusters) + 1}")),
|
||||
description=str(item.get("description", "")),
|
||||
sample_trace_ids=[str(t) for t in item.get("sample_trace_ids", [])],
|
||||
student_failure_rate=failure_rate,
|
||||
teacher_success_rate=success_rate,
|
||||
skill_gap=str(item.get("skill_gap", "")),
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Skipping invalid cluster: %s", e)
|
||||
continue
|
||||
return clusters
|
||||
@@ -0,0 +1,269 @@
|
||||
"""TeacherAgent: frontier model as a tool-calling meta-engineer.
|
||||
|
||||
NOT registered in ``AgentRegistry``. NOT a subclass of ``BaseAgent``.
|
||||
This is a standalone tool-calling loop that wraps ``CloudEngine`` with
|
||||
diagnostic tools for the diagnose phase.
|
||||
|
||||
The teacher:
|
||||
- Uses a frontier model (default ``claude-opus-4-6``) regardless of the
|
||||
user's local intelligence config.
|
||||
- Has its own tool set (diagnostic tools) that user agents do not have.
|
||||
- Tracks cost and stops when the budget is exhausted.
|
||||
- Logs every tool call to a list of ``ToolCallRecord``.
|
||||
|
||||
See spec §5.1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.core.types import Message, Role, ToolCall
|
||||
from openjarvis.learning.distillation.diagnose.types import (
|
||||
DiagnosticTool,
|
||||
ToolCallRecord,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TeacherAgentResult:
|
||||
"""The result of a TeacherAgent.run() call."""
|
||||
|
||||
content: str
|
||||
turns: int
|
||||
total_cost_usd: float
|
||||
total_tokens: int
|
||||
tool_call_records: list[ToolCallRecord] = field(default_factory=list)
|
||||
|
||||
|
||||
class TeacherAgent:
|
||||
"""Frontier model tool-calling loop for the diagnose phase.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
engine :
|
||||
A ``CloudEngine`` (or mock) that provides ``generate()``.
|
||||
model :
|
||||
The frontier model id (e.g. ``"claude-opus-4-6"``).
|
||||
tools :
|
||||
Diagnostic tools exposed to the teacher.
|
||||
max_turns :
|
||||
Maximum number of generate() calls before stopping.
|
||||
max_cost_usd :
|
||||
Maximum accumulated cost before stopping.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
engine: Any,
|
||||
model: str,
|
||||
tools: list[DiagnosticTool],
|
||||
max_turns: int = 30,
|
||||
max_cost_usd: float = 5.0,
|
||||
max_tokens: int = 8192,
|
||||
) -> None:
|
||||
self._engine = engine
|
||||
self._model = model
|
||||
self._tools = {t.name: t for t in tools}
|
||||
self._tool_specs = [t.to_openai_function() for t in tools]
|
||||
self._max_turns = max_turns
|
||||
self._max_cost_usd = max_cost_usd
|
||||
self._max_tokens = max_tokens
|
||||
|
||||
def run(
|
||||
self,
|
||||
user_prompt: str,
|
||||
*,
|
||||
system_prompt: str | None = None,
|
||||
) -> TeacherAgentResult:
|
||||
"""Run the teacher tool-calling loop.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
user_prompt :
|
||||
The instruction to the teacher (e.g. "Analyze student failures").
|
||||
system_prompt :
|
||||
Optional system prompt explaining the teacher's role.
|
||||
|
||||
Returns
|
||||
-------
|
||||
TeacherAgentResult
|
||||
The teacher's final content, cost, and tool call records.
|
||||
"""
|
||||
messages: list[Message] = []
|
||||
if system_prompt:
|
||||
messages.append(Message(role=Role.SYSTEM, content=system_prompt))
|
||||
messages.append(Message(role=Role.USER, content=user_prompt))
|
||||
|
||||
total_cost = 0.0
|
||||
total_tokens = 0
|
||||
tool_call_records: list[ToolCallRecord] = []
|
||||
final_content = ""
|
||||
|
||||
gen_kwargs: dict[str, Any] = {}
|
||||
if self._tool_specs:
|
||||
gen_kwargs["tools"] = self._tool_specs
|
||||
|
||||
for turn in range(1, self._max_turns + 1):
|
||||
# Budget pre-check: don't start a new turn if we've already
|
||||
# exceeded the budget. Without this, the `if not raw_tool_calls`
|
||||
# early-return path below bypasses the post-check entirely,
|
||||
# meaning the final (terminal) turn runs unchecked. Pre-checking
|
||||
# here closes that hole.
|
||||
#
|
||||
# Note: a single turn that ITSELF exceeds the remaining budget
|
||||
# still overshoots. With Opus + ~8k max_tokens + growing input
|
||||
# context, a late-in-conversation final turn can cost $2-4. To
|
||||
# bound this further, lower `max_turns` (caps context growth)
|
||||
# or lower `max_cost_usd` in the config.
|
||||
if total_cost >= self._max_cost_usd:
|
||||
logger.warning(
|
||||
"Teacher budget exceeded before turn %d: $%.2f >= $%.2f",
|
||||
turn,
|
||||
total_cost,
|
||||
self._max_cost_usd,
|
||||
)
|
||||
return TeacherAgentResult(
|
||||
content=final_content,
|
||||
turns=turn - 1,
|
||||
total_cost_usd=total_cost,
|
||||
total_tokens=total_tokens,
|
||||
tool_call_records=tool_call_records,
|
||||
)
|
||||
|
||||
result = self._engine.generate(
|
||||
messages=messages,
|
||||
model=self._model,
|
||||
max_tokens=self._max_tokens,
|
||||
**gen_kwargs,
|
||||
)
|
||||
|
||||
cost = result.get("cost_usd", 0.0)
|
||||
total_cost += cost
|
||||
usage = result.get("usage", {})
|
||||
total_tokens += usage.get("total_tokens", 0)
|
||||
|
||||
content = result.get("content", "")
|
||||
raw_tool_calls = result.get("tool_calls", [])
|
||||
|
||||
if not raw_tool_calls:
|
||||
final_content = content
|
||||
return TeacherAgentResult(
|
||||
content=final_content,
|
||||
turns=turn,
|
||||
total_cost_usd=total_cost,
|
||||
total_tokens=total_tokens,
|
||||
tool_call_records=tool_call_records,
|
||||
)
|
||||
|
||||
# Convert raw tool calls to ToolCall objects
|
||||
tool_call_objs = []
|
||||
for tc in raw_tool_calls:
|
||||
tc_obj = ToolCall(
|
||||
id=tc["id"] if isinstance(tc, dict) else tc.id,
|
||||
name=tc["name"] if isinstance(tc, dict) else tc.name,
|
||||
arguments=(
|
||||
tc.get("arguments", "{}")
|
||||
if isinstance(tc, dict)
|
||||
else tc.arguments
|
||||
),
|
||||
)
|
||||
tool_call_objs.append(tc_obj)
|
||||
|
||||
# Append assistant message with tool calls
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.ASSISTANT,
|
||||
content=content,
|
||||
tool_calls=tool_call_objs,
|
||||
)
|
||||
)
|
||||
|
||||
# Execute each tool call
|
||||
for tc in tool_call_objs:
|
||||
tc_name = tc.name
|
||||
tc_id = tc.id
|
||||
tc_args_str = tc.arguments
|
||||
|
||||
try:
|
||||
tc_args = json.loads(tc_args_str)
|
||||
except json.JSONDecodeError:
|
||||
tc_args = {}
|
||||
|
||||
tool = self._tools.get(tc_name)
|
||||
start_time = time.monotonic()
|
||||
if tool is not None:
|
||||
try:
|
||||
tool_result = tool.fn(**tc_args)
|
||||
except Exception as e:
|
||||
tool_result = json.dumps({"error": str(e)})
|
||||
logger.warning("Tool %s raised: %s", tc_name, e)
|
||||
else:
|
||||
tool_result = json.dumps({"error": f"Unknown tool: {tc_name}"})
|
||||
elapsed_ms = (time.monotonic() - start_time) * 1000
|
||||
|
||||
# Safety cap on tool results to bound context growth.
|
||||
# Primary truncation should happen inside each tool, but this
|
||||
# is defense-in-depth for tools that return too much. 20KB
|
||||
# per call × 30 max_turns = 600KB total, well under Opus's
|
||||
# 1M-token context window.
|
||||
MAX_TOOL_RESULT_CHARS = 20000
|
||||
if len(tool_result) > MAX_TOOL_RESULT_CHARS:
|
||||
tool_result = (
|
||||
tool_result[:MAX_TOOL_RESULT_CHARS]
|
||||
+ f"\n...[tool result truncated from {len(tool_result)} chars]"
|
||||
)
|
||||
|
||||
tool_call_records.append(
|
||||
ToolCallRecord(
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
tool=tc_name,
|
||||
args=tc_args,
|
||||
result=tool_result[:8000], # Shorter cap for audit log
|
||||
latency_ms=elapsed_ms,
|
||||
cost_usd=0.0, # Tool calls themselves are free
|
||||
)
|
||||
)
|
||||
|
||||
# Append tool result message (uses the safety-capped value)
|
||||
messages.append(
|
||||
Message(
|
||||
role=Role.TOOL,
|
||||
content=tool_result,
|
||||
tool_call_id=tc_id,
|
||||
name=tc_name,
|
||||
)
|
||||
)
|
||||
|
||||
# Check cost budget
|
||||
if total_cost >= self._max_cost_usd:
|
||||
logger.warning(
|
||||
"Teacher cost budget exceeded: $%.2f >= $%.2f",
|
||||
total_cost,
|
||||
self._max_cost_usd,
|
||||
)
|
||||
final_content = content
|
||||
return TeacherAgentResult(
|
||||
content=final_content,
|
||||
turns=turn,
|
||||
total_cost_usd=total_cost,
|
||||
total_tokens=total_tokens,
|
||||
tool_call_records=tool_call_records,
|
||||
)
|
||||
|
||||
# Exhausted max_turns
|
||||
logger.warning("Teacher exhausted max_turns=%d", self._max_turns)
|
||||
return TeacherAgentResult(
|
||||
content=final_content,
|
||||
turns=self._max_turns,
|
||||
total_cost_usd=total_cost,
|
||||
total_tokens=total_tokens,
|
||||
tool_call_records=tool_call_records,
|
||||
)
|
||||
@@ -0,0 +1,543 @@
|
||||
"""Diagnostic tools exposed to the teacher in the diagnose phase.
|
||||
|
||||
All tools are **read-only** relative to the user's config. They do not mutate
|
||||
``~/.openjarvis/config.toml``, agent prompts, or tool descriptions.
|
||||
|
||||
Tools that execute code (``run_student_on_task``, ``run_self_on_task``) append
|
||||
new traces to ``TraceStore`` as a side effect. These traces are tagged with
|
||||
``source=distillation_session:<id>`` so they can be excluded from future
|
||||
learning input.
|
||||
|
||||
See spec §5.2.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.learning.distillation.diagnose.types import DiagnosticTool
|
||||
|
||||
|
||||
def build_diagnostic_tools(
|
||||
*,
|
||||
trace_store: Any,
|
||||
config: dict[str, Any],
|
||||
benchmark_samples: list,
|
||||
student_runner: Any,
|
||||
teacher_engine: Any,
|
||||
teacher_model: str,
|
||||
judge: Any,
|
||||
session_id: str,
|
||||
) -> list[DiagnosticTool]:
|
||||
"""Build all diagnostic tools as closures over shared dependencies.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
trace_store :
|
||||
A ``TraceStore`` instance for trace queries.
|
||||
config :
|
||||
Dict with ``config_path`` (Path) and ``openjarvis_home`` (Path).
|
||||
benchmark_samples :
|
||||
List of ``PersonalBenchmarkSample`` objects.
|
||||
student_runner :
|
||||
Callable that re-executes the student on a task.
|
||||
teacher_engine :
|
||||
The ``CloudEngine`` instance used by the teacher.
|
||||
teacher_model :
|
||||
Model id for teacher inference (e.g. "claude-opus-4-6").
|
||||
judge :
|
||||
A ``TraceJudge`` instance for comparing outputs.
|
||||
session_id :
|
||||
Current session id for tagging traces.
|
||||
"""
|
||||
openjarvis_home = config["openjarvis_home"]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# list_traces
|
||||
# ------------------------------------------------------------------
|
||||
def _list_traces(
|
||||
limit: int = 20,
|
||||
agent: str | None = None,
|
||||
outcome: str | None = None,
|
||||
min_feedback: float | None = None,
|
||||
max_feedback: float | None = None,
|
||||
) -> str:
|
||||
kwargs: dict[str, Any] = {"limit": limit}
|
||||
if agent:
|
||||
kwargs["agent"] = agent
|
||||
if outcome:
|
||||
kwargs["outcome"] = outcome
|
||||
traces = trace_store.list_traces(**kwargs)
|
||||
metas = []
|
||||
for t in traces:
|
||||
fb = getattr(t, "feedback", None)
|
||||
if min_feedback is not None and (fb is None or fb < min_feedback):
|
||||
continue
|
||||
if max_feedback is not None and (fb is None or fb > max_feedback):
|
||||
continue
|
||||
metas.append(
|
||||
{
|
||||
"trace_id": t.trace_id,
|
||||
"query": t.query[:200],
|
||||
"agent": t.agent,
|
||||
"model": t.model,
|
||||
"outcome": t.outcome,
|
||||
"feedback": fb,
|
||||
"started_at": t.started_at,
|
||||
}
|
||||
)
|
||||
return json.dumps(metas, default=str)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# get_trace
|
||||
# ------------------------------------------------------------------
|
||||
def _get_trace(trace_id: str) -> str:
|
||||
trace = trace_store.get(trace_id)
|
||||
if trace is None:
|
||||
return json.dumps({"error": f"Trace {trace_id} not found"})
|
||||
# Truncate large free-text fields so a single get_trace call can't
|
||||
# blow out the teacher's context window. monitor_operative traces
|
||||
# can have 50-95K tokens of result text each.
|
||||
query = trace.query or ""
|
||||
result = trace.result or ""
|
||||
max_query, max_result = 2000, 6000
|
||||
query_display = query[:max_query] + (
|
||||
f"...[query truncated from {len(query)} chars]"
|
||||
if len(query) > max_query
|
||||
else ""
|
||||
)
|
||||
result_display = result[:max_result] + (
|
||||
f"...[result truncated from {len(result)} chars]"
|
||||
if len(result) > max_result
|
||||
else ""
|
||||
)
|
||||
return json.dumps(
|
||||
{
|
||||
"trace_id": trace.trace_id,
|
||||
"query": query_display,
|
||||
"agent": trace.agent,
|
||||
"model": trace.model,
|
||||
"outcome": trace.outcome,
|
||||
"feedback": getattr(trace, "feedback", None),
|
||||
"result": result_display,
|
||||
"total_tokens": trace.total_tokens,
|
||||
"total_latency_seconds": trace.total_latency_seconds,
|
||||
"steps": [str(s)[:200] for s in getattr(trace, "steps", [])[:20]],
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# search_traces
|
||||
# ------------------------------------------------------------------
|
||||
def _search_traces(query: str, limit: int = 20) -> str:
|
||||
# Cap limit to avoid pathological calls that would OOM the context.
|
||||
# 20 x 95K-token traces = ~1.9M tokens, well over Opus's 1M window.
|
||||
limit = max(1, min(limit, 20))
|
||||
try:
|
||||
results = trace_store.search(query, limit=limit)
|
||||
except Exception as e:
|
||||
# SQLite FTS5 raises "unknown special query" for patterns with
|
||||
# unescaped special chars (- / . etc). Return a structured error
|
||||
# so the teacher can retry with a cleaner query instead of the
|
||||
# call being swallowed and context still growing.
|
||||
return json.dumps(
|
||||
{
|
||||
"error": f"search_traces failed: {e}",
|
||||
"hint": (
|
||||
"FTS5 special chars (- / . : etc) need escaping. "
|
||||
"Try a simpler keyword query."
|
||||
),
|
||||
}
|
||||
)
|
||||
# Truncate per-trace text fields. Each trace's `result` column can be
|
||||
# 50-95K tokens; returning raw content would blow the context window.
|
||||
max_query_per_trace = 300
|
||||
max_result_per_trace = 1500
|
||||
for r in results:
|
||||
q = r.get("query") or ""
|
||||
if len(q) > max_query_per_trace:
|
||||
r["query"] = q[:max_query_per_trace] + f"...[{len(q)} chars]"
|
||||
res = r.get("result") or ""
|
||||
if len(res) > max_result_per_trace:
|
||||
r["result"] = res[:max_result_per_trace] + f"...[{len(res)} chars]"
|
||||
return json.dumps(results, default=str)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# get_current_config
|
||||
# ------------------------------------------------------------------
|
||||
def _get_current_config() -> str:
|
||||
config_path = config["config_path"]
|
||||
try:
|
||||
return config_path.read_text(encoding="utf-8")
|
||||
except FileNotFoundError:
|
||||
return "No config.toml found."
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# get_agent_prompt
|
||||
# ------------------------------------------------------------------
|
||||
def _get_agent_prompt(agent_name: str) -> str:
|
||||
prompt_path = openjarvis_home / "agents" / agent_name / "system_prompt.md"
|
||||
try:
|
||||
return prompt_path.read_text(encoding="utf-8")
|
||||
except FileNotFoundError:
|
||||
return f"No system prompt found for agent '{agent_name}'."
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# get_tool_description
|
||||
# ------------------------------------------------------------------
|
||||
def _get_tool_description(tool_name: str) -> str:
|
||||
desc_path = openjarvis_home / "tools" / "descriptions.toml"
|
||||
try:
|
||||
content = desc_path.read_text(encoding="utf-8")
|
||||
# Simple TOML parsing for the description field
|
||||
in_section = False
|
||||
for line in content.splitlines():
|
||||
if line.strip() == f"[{tool_name}]":
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.startswith("["):
|
||||
break
|
||||
if in_section and "description" in line:
|
||||
return line.split("=", 1)[1].strip().strip('"')
|
||||
return f"Tool '{tool_name}' found but no description field."
|
||||
except FileNotFoundError:
|
||||
return "No descriptions.toml found."
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# list_available_tools
|
||||
# ------------------------------------------------------------------
|
||||
def _list_available_tools() -> str:
|
||||
# Read from the on-disk descriptions.toml
|
||||
desc_path = openjarvis_home / "tools" / "descriptions.toml"
|
||||
tools_list = []
|
||||
try:
|
||||
content = desc_path.read_text(encoding="utf-8")
|
||||
current_tool = None
|
||||
for line in content.splitlines():
|
||||
if line.startswith("[") and line.endswith("]"):
|
||||
current_tool = line[1:-1]
|
||||
tools_list.append(
|
||||
{
|
||||
"name": current_tool,
|
||||
"description": "",
|
||||
"category": "general",
|
||||
"agents": [],
|
||||
}
|
||||
)
|
||||
elif current_tool and "description" in line and "=" in line:
|
||||
desc = line.split("=", 1)[1].strip().strip('"')
|
||||
if tools_list:
|
||||
tools_list[-1]["description"] = desc
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
return json.dumps(tools_list, default=str)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# list_personal_benchmark
|
||||
# ------------------------------------------------------------------
|
||||
def _list_personal_benchmark(limit: int = 50) -> str:
|
||||
tasks = []
|
||||
for sample in benchmark_samples[:limit]:
|
||||
tasks.append(
|
||||
{
|
||||
"task_id": sample.trace_id,
|
||||
"query": sample.query,
|
||||
"reference_answer": sample.reference_answer[:500],
|
||||
"category": getattr(sample, "category", "chat"),
|
||||
}
|
||||
)
|
||||
return json.dumps(tasks, default=str)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# run_student_on_task
|
||||
# ------------------------------------------------------------------
|
||||
def _run_student_on_task(task_id: str) -> str:
|
||||
sample = next((s for s in benchmark_samples if s.trace_id == task_id), None)
|
||||
if sample is None:
|
||||
return json.dumps({"error": f"Task {task_id} not found in benchmark"})
|
||||
result = student_runner(sample.query, session_id=session_id)
|
||||
return json.dumps(
|
||||
{
|
||||
"task_id": task_id,
|
||||
"output": str(getattr(result, "content", result)),
|
||||
"score": getattr(result, "score", 0.0),
|
||||
"trace_id": getattr(result, "trace_id", ""),
|
||||
"latency_seconds": getattr(result, "latency_seconds", 0.0),
|
||||
"tokens_used": getattr(result, "tokens_used", 0),
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# run_self_on_task
|
||||
# ------------------------------------------------------------------
|
||||
def _run_self_on_task(task_id: str, max_tokens: int = 2048) -> str:
|
||||
sample = next((s for s in benchmark_samples if s.trace_id == task_id), None)
|
||||
if sample is None:
|
||||
return json.dumps({"error": f"Task {task_id} not found in benchmark"})
|
||||
response = teacher_engine.generate(
|
||||
messages=[{"role": "user", "content": sample.query}],
|
||||
model=teacher_model,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
return json.dumps(
|
||||
{
|
||||
"task_id": task_id,
|
||||
"output": response.get("content", ""),
|
||||
"reasoning": "",
|
||||
"cost_usd": response.get("cost_usd", 0.0),
|
||||
"tokens_used": response.get("usage", {}).get("total_tokens", 0),
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# compare_outputs
|
||||
# ------------------------------------------------------------------
|
||||
def _compare_outputs(student_output: str, teacher_output: str, task: str) -> str:
|
||||
score, reasoning = judge.score_trace(
|
||||
type(
|
||||
"FakeTrace",
|
||||
(),
|
||||
{
|
||||
"query": task,
|
||||
"result": student_output,
|
||||
"steps": [],
|
||||
"messages": [],
|
||||
"agent": "student",
|
||||
"model": "local",
|
||||
"total_tokens": 0,
|
||||
"total_latency_seconds": 0,
|
||||
"metadata": {},
|
||||
},
|
||||
)()
|
||||
)
|
||||
teacher_score, _ = judge.score_trace(
|
||||
type(
|
||||
"FakeTrace",
|
||||
(),
|
||||
{
|
||||
"query": task,
|
||||
"result": teacher_output,
|
||||
"steps": [],
|
||||
"messages": [],
|
||||
"agent": "teacher",
|
||||
"model": "frontier",
|
||||
"total_tokens": 0,
|
||||
"total_latency_seconds": 0,
|
||||
"metadata": {},
|
||||
},
|
||||
)()
|
||||
)
|
||||
return json.dumps(
|
||||
{
|
||||
"task_id": "",
|
||||
"student_score": score,
|
||||
"teacher_score": teacher_score,
|
||||
"judge_reasoning": reasoning,
|
||||
},
|
||||
default=str,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Assemble
|
||||
# ------------------------------------------------------------------
|
||||
return [
|
||||
DiagnosticTool(
|
||||
name="list_traces",
|
||||
description=(
|
||||
"Browse traces by agent, outcome, feedback range."
|
||||
" Returns a JSON list of trace summaries."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Max results",
|
||||
"default": 20,
|
||||
},
|
||||
"agent": {
|
||||
"type": "string",
|
||||
"description": "Filter by agent name",
|
||||
},
|
||||
"outcome": {
|
||||
"type": "string",
|
||||
"description": "Filter by outcome: success/failure",
|
||||
},
|
||||
"min_feedback": {
|
||||
"type": "number",
|
||||
"description": "Minimum feedback score",
|
||||
},
|
||||
"max_feedback": {
|
||||
"type": "number",
|
||||
"description": "Maximum feedback score",
|
||||
},
|
||||
},
|
||||
},
|
||||
fn=_list_traces,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="get_trace",
|
||||
description=(
|
||||
"Read a single trace including query, result, steps, and metrics."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"trace_id": {
|
||||
"type": "string",
|
||||
"description": "The trace ID to retrieve",
|
||||
},
|
||||
},
|
||||
"required": ["trace_id"],
|
||||
},
|
||||
fn=_get_trace,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="search_traces",
|
||||
description=(
|
||||
"Full-text search across traces. Returns matching trace summaries."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": {"type": "string", "description": "Search query"},
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Max results",
|
||||
"default": 20,
|
||||
},
|
||||
},
|
||||
"required": ["query"],
|
||||
},
|
||||
fn=_search_traces,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="get_current_config",
|
||||
description="Read the current OpenJarvis config.toml.",
|
||||
parameters={"type": "object", "properties": {}},
|
||||
fn=_get_current_config,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="get_agent_prompt",
|
||||
description="Read the current system prompt for a named agent.",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"agent_name": {
|
||||
"type": "string",
|
||||
"description": "Agent name (e.g. 'simple')",
|
||||
},
|
||||
},
|
||||
"required": ["agent_name"],
|
||||
},
|
||||
fn=_get_agent_prompt,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="get_tool_description",
|
||||
description="Read the LM-facing description of a tool.",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"tool_name": {
|
||||
"type": "string",
|
||||
"description": "Tool name (e.g. 'web_search')",
|
||||
},
|
||||
},
|
||||
"required": ["tool_name"],
|
||||
},
|
||||
fn=_get_tool_description,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="list_available_tools",
|
||||
description="List all tools and which agents currently have them enabled.",
|
||||
parameters={"type": "object", "properties": {}},
|
||||
fn=_list_available_tools,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="list_personal_benchmark",
|
||||
description=(
|
||||
"Browse personal benchmark tasks with their reference answers."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"limit": {
|
||||
"type": "integer",
|
||||
"description": "Max tasks to return",
|
||||
"default": 50,
|
||||
},
|
||||
},
|
||||
},
|
||||
fn=_list_personal_benchmark,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="run_student_on_task",
|
||||
description=(
|
||||
"Re-execute the local student agent on a benchmark task."
|
||||
" Returns the student's output and score."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"task_id": {
|
||||
"type": "string",
|
||||
"description": "Benchmark task ID",
|
||||
},
|
||||
},
|
||||
"required": ["task_id"],
|
||||
},
|
||||
fn=_run_student_on_task,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="run_self_on_task",
|
||||
description=(
|
||||
"Run yourself (the teacher) on a benchmark task"
|
||||
" to produce a reference answer."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"task_id": {
|
||||
"type": "string",
|
||||
"description": "Benchmark task ID",
|
||||
},
|
||||
"max_tokens": {
|
||||
"type": "integer",
|
||||
"description": "Max tokens for response",
|
||||
"default": 2048,
|
||||
},
|
||||
},
|
||||
"required": ["task_id"],
|
||||
},
|
||||
fn=_run_self_on_task,
|
||||
),
|
||||
DiagnosticTool(
|
||||
name="compare_outputs",
|
||||
description=(
|
||||
"Compare student and teacher outputs on a task using the TraceJudge."
|
||||
),
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"student_output": {
|
||||
"type": "string",
|
||||
"description": "The student's output",
|
||||
},
|
||||
"teacher_output": {
|
||||
"type": "string",
|
||||
"description": "The teacher's output",
|
||||
},
|
||||
"task": {
|
||||
"type": "string",
|
||||
"description": "The original task/query",
|
||||
},
|
||||
},
|
||||
"required": ["student_output", "teacher_output", "task"],
|
||||
},
|
||||
fn=_compare_outputs,
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1,129 @@
|
||||
"""Data types for the diagnose phase.
|
||||
|
||||
Lightweight dataclasses used as return types by diagnostic tools and
|
||||
as internal data carriers. These are NOT pydantic models — they don't
|
||||
need validation or JSON schema generation.
|
||||
|
||||
See spec §5.2 for the tool return type rationale.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TraceMeta:
|
||||
"""Lightweight summary of a trace for browsing."""
|
||||
|
||||
trace_id: str
|
||||
query: str
|
||||
agent: str
|
||||
model: str
|
||||
outcome: Optional[str]
|
||||
feedback: Optional[float]
|
||||
started_at: float
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class BenchmarkTask:
|
||||
"""One task from the personal benchmark."""
|
||||
|
||||
task_id: str
|
||||
query: str
|
||||
reference_answer: str
|
||||
category: str = "chat"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class StudentRun:
|
||||
"""Result of re-executing the local student on a benchmark task."""
|
||||
|
||||
task_id: str
|
||||
output: str
|
||||
score: float
|
||||
trace_id: str
|
||||
latency_seconds: float
|
||||
tokens_used: int
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TeacherRun:
|
||||
"""Result of the teacher running itself on a benchmark task."""
|
||||
|
||||
task_id: str
|
||||
output: str
|
||||
reasoning: str
|
||||
cost_usd: float
|
||||
tokens_used: int
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ComparisonResult:
|
||||
"""Structured comparison between student and teacher outputs."""
|
||||
|
||||
task_id: str
|
||||
student_score: float
|
||||
teacher_score: float
|
||||
judge_reasoning: str
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ToolMeta:
|
||||
"""Metadata about a tool in the ToolRegistry."""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
category: str
|
||||
agents: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DiagnosticTool:
|
||||
"""A tool exposed to the teacher in the diagnose phase.
|
||||
|
||||
Unlike ``BaseTool``, these are not registered in ``ToolRegistry``.
|
||||
They are lightweight wrappers: a name, description, JSON schema
|
||||
for parameters, and a callable that implements the tool.
|
||||
"""
|
||||
|
||||
name: str
|
||||
description: str
|
||||
parameters: dict[str, Any]
|
||||
fn: Callable[..., Any]
|
||||
|
||||
def to_openai_function(self) -> dict[str, Any]:
|
||||
"""Convert to OpenAI function-calling format."""
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"parameters": self.parameters,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ToolCallRecord:
|
||||
"""One teacher tool call, persisted to the JSONL log."""
|
||||
|
||||
timestamp: datetime
|
||||
tool: str
|
||||
args: dict[str, Any]
|
||||
result: str
|
||||
latency_ms: float
|
||||
cost_usd: float
|
||||
|
||||
def to_jsonl_dict(self) -> dict[str, Any]:
|
||||
"""Serialize to a JSON-safe dict for JSONL output."""
|
||||
return {
|
||||
"timestamp": self.timestamp.isoformat(),
|
||||
"tool": self.tool,
|
||||
"args": self.args,
|
||||
"result": self.result,
|
||||
"latency_ms": self.latency_ms,
|
||||
"cost_usd": self.cost_usd,
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
"""Execute phase: apply edits to the harness configuration."""
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Concrete EditApplier implementations.
|
||||
|
||||
Importing this package triggers registration of all appliers
|
||||
in the EditApplierRegistry. Applier modules are imported in Task 6
|
||||
once all implementations exist.
|
||||
"""
|
||||
@@ -0,0 +1,226 @@
|
||||
"""Agent-pillar appliers: prompts, class, params, few-shot.
|
||||
|
||||
See spec §4.1 op semantics for agent ops.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
ApplyResult,
|
||||
EditApplier,
|
||||
ValidationResult,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import apply_unified_diff
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _agent_prompt_path(ctx: ApplyContext, agent_name: str) -> Path:
|
||||
return ctx.agents_dir / agent_name / "system_prompt.md"
|
||||
|
||||
|
||||
def _extract_agent_name(edit: Edit) -> str:
|
||||
"""Extract agent name from edit target or payload."""
|
||||
if "agent" in edit.payload:
|
||||
return edit.payload["agent"]
|
||||
# Try from target: "agents.simple.system_prompt" -> "simple"
|
||||
parts = edit.target.split(".")
|
||||
if len(parts) >= 2:
|
||||
return parts[1]
|
||||
return "default"
|
||||
|
||||
|
||||
class ReplaceSystemPromptApplier(EditApplier):
|
||||
"""Overwrite an agent's entire system prompt."""
|
||||
|
||||
op = EditOp.REPLACE_SYSTEM_PROMPT
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
if "new_content" not in edit.payload:
|
||||
return ValidationResult(ok=False, reason="Missing new_content in payload")
|
||||
agent = _extract_agent_name(edit)
|
||||
prompt_path = _agent_prompt_path(ctx, agent)
|
||||
if not prompt_path.parent.exists():
|
||||
return ValidationResult(
|
||||
ok=False, reason=f"Agent directory not found: {prompt_path.parent}"
|
||||
)
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = _extract_agent_name(edit)
|
||||
prompt_path = _agent_prompt_path(ctx, agent)
|
||||
prompt_path.write_text(edit.payload["new_content"], encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(prompt_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class PatchSystemPromptApplier(EditApplier):
|
||||
"""Apply a unified diff to an agent's system prompt."""
|
||||
|
||||
op = EditOp.PATCH_SYSTEM_PROMPT
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
if "diff" not in edit.payload:
|
||||
return ValidationResult(ok=False, reason="Missing diff in payload")
|
||||
agent = _extract_agent_name(edit)
|
||||
prompt_path = _agent_prompt_path(ctx, agent)
|
||||
if not prompt_path.exists():
|
||||
return ValidationResult(
|
||||
ok=False, reason=f"Prompt file not found: {prompt_path}"
|
||||
)
|
||||
original = prompt_path.read_text(encoding="utf-8")
|
||||
patched = apply_unified_diff(original, edit.payload["diff"])
|
||||
if patched is None:
|
||||
return ValidationResult(
|
||||
ok=False, reason="Diff cannot be applied to current prompt"
|
||||
)
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = _extract_agent_name(edit)
|
||||
prompt_path = _agent_prompt_path(ctx, agent)
|
||||
original = prompt_path.read_text(encoding="utf-8")
|
||||
patched = apply_unified_diff(original, edit.payload["diff"])
|
||||
if patched is None:
|
||||
raise RuntimeError(f"Failed to apply diff to {prompt_path}")
|
||||
prompt_path.write_text(patched, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(prompt_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class SetAgentClassApplier(EditApplier):
|
||||
"""Change which agent class is used."""
|
||||
|
||||
op = EditOp.SET_AGENT_CLASS
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
if "new_class" not in edit.payload:
|
||||
return ValidationResult(ok=False, reason="Missing new_class in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = _extract_agent_name(edit)
|
||||
new_class = edit.payload["new_class"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
) # noqa: E501
|
||||
|
||||
section = f"[agent.{agent}]"
|
||||
if section in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
found = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
break
|
||||
if in_section and line.strip().startswith("class"):
|
||||
lines[i] = f'class = "{new_class}"'
|
||||
found = True
|
||||
break
|
||||
if not found:
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
lines.insert(i + 1, f'class = "{new_class}"')
|
||||
break
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
content += f'\n{section}\nclass = "{new_class}"\n'
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class SetAgentParamApplier(EditApplier):
|
||||
"""Update an agent parameter."""
|
||||
|
||||
op = EditOp.SET_AGENT_PARAM
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
for key in ("agent", "param", "value"):
|
||||
if key not in edit.payload:
|
||||
return ValidationResult(ok=False, reason=f"Missing {key} in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = edit.payload["agent"]
|
||||
param = edit.payload["param"]
|
||||
value = edit.payload["value"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
) # noqa: E501
|
||||
|
||||
section = f"[agent.{agent}]"
|
||||
if section in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
found = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
lines.insert(i, f"{param} = {value}")
|
||||
found = True
|
||||
break
|
||||
if in_section and line.strip().startswith(f"{param}"):
|
||||
lines[i] = f"{param} = {value}"
|
||||
found = True
|
||||
break
|
||||
if not found:
|
||||
lines.append(f"{param} = {value}")
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
content += f"\n{section}\n{param} = {value}\n"
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class EditFewShotExemplarsApplier(EditApplier):
|
||||
"""Write few-shot exemplars to an agent's directory."""
|
||||
|
||||
op = EditOp.EDIT_FEW_SHOT_EXEMPLARS
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
if "exemplars" not in edit.payload:
|
||||
return ValidationResult(ok=False, reason="Missing exemplars in payload")
|
||||
agent = _extract_agent_name(edit)
|
||||
agent_dir = ctx.agents_dir / agent
|
||||
if not agent_dir.exists():
|
||||
return ValidationResult(
|
||||
ok=False, reason=f"Agent directory not found: {agent_dir}"
|
||||
)
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = _extract_agent_name(edit)
|
||||
fs_path = ctx.agents_dir / agent / "few_shot.json"
|
||||
fs_path.write_text(
|
||||
json.dumps(edit.payload["exemplars"], indent=2),
|
||||
encoding="utf-8",
|
||||
)
|
||||
return ApplyResult(changed_files=[str(fs_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
@@ -0,0 +1,117 @@
|
||||
"""Intelligence-pillar appliers: model routing and parameters.
|
||||
|
||||
See spec §4.1 op semantics for SET_MODEL_FOR_QUERY_CLASS and SET_MODEL_PARAM.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
ApplyResult,
|
||||
EditApplier,
|
||||
ValidationResult,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
|
||||
class SetModelForQueryClassApplier(EditApplier):
|
||||
"""Update the routing policy map for a query class."""
|
||||
|
||||
op = EditOp.SET_MODEL_FOR_QUERY_CLASS
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
if "query_class" not in edit.payload or "model" not in edit.payload:
|
||||
return ValidationResult(
|
||||
ok=False, reason="Missing query_class or model in payload"
|
||||
)
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
query_class = edit.payload["query_class"]
|
||||
model = edit.payload["model"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
)
|
||||
|
||||
# Check if the policy_map section exists
|
||||
if "[learning.routing.policy_map]" in content:
|
||||
# Check if this query_class already has a line
|
||||
lines = content.splitlines()
|
||||
found = False
|
||||
for i, line in enumerate(lines):
|
||||
stripped = line.strip()
|
||||
if stripped.startswith(f"{query_class}") and "=" in stripped:
|
||||
lines[i] = f'{query_class} = "{model}"'
|
||||
found = True
|
||||
break
|
||||
if not found:
|
||||
# Append after the [learning.routing.policy_map] header
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == "[learning.routing.policy_map]":
|
||||
lines.insert(i + 1, f'{query_class} = "{model}"')
|
||||
break
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
# Append the section
|
||||
content += f'\n[learning.routing.policy_map]\n{query_class} = "{model}"\n'
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass # Delegated to CheckpointStore.discard_stage() in the execution loop
|
||||
|
||||
|
||||
class SetModelParamApplier(EditApplier):
|
||||
"""Update a model parameter in config."""
|
||||
|
||||
op = EditOp.SET_MODEL_PARAM
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
for key in ("model", "param", "value"):
|
||||
if key not in edit.payload:
|
||||
return ValidationResult(ok=False, reason=f"Missing {key} in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
model = edit.payload["model"]
|
||||
param = edit.payload["param"]
|
||||
value = edit.payload["value"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
)
|
||||
|
||||
# Sanitize model name for TOML section header (replace : with -)
|
||||
section_key = model.replace(":", "-")
|
||||
section_header = f"[models.{section_key}]"
|
||||
|
||||
if section_header in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
found = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section_header:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
# Hit next section — insert before it
|
||||
lines.insert(i, f"{param} = {value}")
|
||||
found = True
|
||||
break
|
||||
if in_section and line.strip().startswith(f"{param}"):
|
||||
lines[i] = f"{param} = {value}"
|
||||
found = True
|
||||
break
|
||||
if not found:
|
||||
lines.append(f"{param} = {value}")
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
content += f"\n{section_header}\n{param} = {value}\n"
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass # Delegated to CheckpointStore.discard_stage()
|
||||
@@ -0,0 +1,38 @@
|
||||
"""LoRA fine-tuning stub — deferred to v2.
|
||||
|
||||
The planner can emit LORA_FINETUNE edits so the diagnosis surfaces
|
||||
"this should be a weight update" pressure to the user, but the executor
|
||||
refuses them in v1.
|
||||
|
||||
See spec §4.1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
ApplyResult,
|
||||
EditApplier,
|
||||
ValidationResult,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
|
||||
class LoraStubApplier(EditApplier):
|
||||
"""Refuses LORA_FINETUNE with a clear v2 message."""
|
||||
|
||||
op = EditOp.LORA_FINETUNE
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
return ValidationResult(
|
||||
ok=False,
|
||||
reason="LORA_FINETUNE is deferred to v2. "
|
||||
"The planner emitted this edit to signal that weight updates "
|
||||
"would help, but the executor cannot apply them yet.",
|
||||
)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
raise NotImplementedError("LORA_FINETUNE deferred to v2")
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
@@ -0,0 +1,159 @@
|
||||
"""Tools-pillar appliers: add/remove tools, edit descriptions.
|
||||
|
||||
See spec §4.1 op semantics for tool ops.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
ApplyResult,
|
||||
EditApplier,
|
||||
ValidationResult,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
|
||||
class AddToolToAgentApplier(EditApplier):
|
||||
"""Add a tool to an agent's tool list in config."""
|
||||
|
||||
op = EditOp.ADD_TOOL_TO_AGENT
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
for key in ("agent", "tool_name"):
|
||||
if key not in edit.payload:
|
||||
return ValidationResult(ok=False, reason=f"Missing {key} in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = edit.payload["agent"]
|
||||
tool_name = edit.payload["tool_name"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
)
|
||||
|
||||
section = f"[agent.{agent}]"
|
||||
if section in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
break
|
||||
if in_section and line.strip().startswith("tools"):
|
||||
# Parse existing tool list and add new tool
|
||||
match = re.search(r"\[([^\]]*)\]", line)
|
||||
if match:
|
||||
existing = match.group(1)
|
||||
if tool_name not in existing:
|
||||
new_list = existing.rstrip() + f', "{tool_name}"'
|
||||
lines[i] = f"tools = [{new_list}]"
|
||||
break
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
content += f'\n{section}\ntools = ["{tool_name}"]\n'
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class RemoveToolFromAgentApplier(EditApplier):
|
||||
"""Remove a tool from an agent's tool list in config."""
|
||||
|
||||
op = EditOp.REMOVE_TOOL_FROM_AGENT
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
for key in ("agent", "tool_name"):
|
||||
if key not in edit.payload:
|
||||
return ValidationResult(ok=False, reason=f"Missing {key} in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
agent = edit.payload["agent"]
|
||||
tool_name = edit.payload["tool_name"]
|
||||
config_path = ctx.config_path
|
||||
content = (
|
||||
config_path.read_text(encoding="utf-8") if config_path.exists() else ""
|
||||
)
|
||||
|
||||
section = f"[agent.{agent}]"
|
||||
if section in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
break
|
||||
if in_section and line.strip().startswith("tools"):
|
||||
# Remove the tool from the list
|
||||
line = re.sub(rf',?\s*"{re.escape(tool_name)}"', "", line)
|
||||
line = re.sub(rf'"{re.escape(tool_name)}"\s*,?\s*', "", line)
|
||||
lines[i] = line
|
||||
break
|
||||
content = "\n".join(lines) + "\n"
|
||||
|
||||
config_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(config_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class EditToolDescriptionApplier(EditApplier):
|
||||
"""Update a tool's LM-facing description in descriptions.toml."""
|
||||
|
||||
op = EditOp.EDIT_TOOL_DESCRIPTION
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
for key in ("tool_name", "new_description"):
|
||||
if key not in edit.payload:
|
||||
return ValidationResult(ok=False, reason=f"Missing {key} in payload")
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
tool_name = edit.payload["tool_name"]
|
||||
new_desc = edit.payload["new_description"]
|
||||
desc_path = ctx.tools_dir / "descriptions.toml"
|
||||
desc_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
content = desc_path.read_text(encoding="utf-8") if desc_path.exists() else ""
|
||||
|
||||
section = f"[{tool_name}]"
|
||||
if section in content:
|
||||
lines = content.splitlines()
|
||||
in_section = False
|
||||
found = False
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
in_section = True
|
||||
continue
|
||||
if in_section and line.strip().startswith("["):
|
||||
break
|
||||
if in_section and line.strip().startswith("description"):
|
||||
lines[i] = f'description = "{new_desc}"'
|
||||
found = True
|
||||
break
|
||||
if not found:
|
||||
for i, line in enumerate(lines):
|
||||
if line.strip() == section:
|
||||
lines.insert(i + 1, f'description = "{new_desc}"')
|
||||
break
|
||||
content = "\n".join(lines) + "\n"
|
||||
else:
|
||||
content += f'\n{section}\ndescription = "{new_desc}"\n'
|
||||
|
||||
desc_path.write_text(content, encoding="utf-8")
|
||||
return ApplyResult(changed_files=[str(desc_path)])
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
@@ -0,0 +1,95 @@
|
||||
"""EditApplier ABC, registry, and context types for the execute phase.
|
||||
|
||||
Each concrete applier implements validate/apply/rollback for a single EditOp.
|
||||
Appliers are registered in an EditApplierRegistry keyed by EditOp.
|
||||
|
||||
See spec §7.1.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import ClassVar
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
|
||||
@dataclass
|
||||
class ApplyContext:
|
||||
"""Shared context passed to all appliers."""
|
||||
|
||||
openjarvis_home: Path
|
||||
session_id: str
|
||||
|
||||
@property
|
||||
def config_path(self) -> Path:
|
||||
return self.openjarvis_home / "config.toml"
|
||||
|
||||
@property
|
||||
def agents_dir(self) -> Path:
|
||||
return self.openjarvis_home / "agents"
|
||||
|
||||
@property
|
||||
def tools_dir(self) -> Path:
|
||||
return self.openjarvis_home / "tools"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidationResult:
|
||||
"""Result of EditApplier.validate()."""
|
||||
|
||||
ok: bool
|
||||
reason: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class ApplyResult:
|
||||
"""Result of EditApplier.apply()."""
|
||||
|
||||
changed_files: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
class EditApplier(ABC):
|
||||
"""Abstract base for edit appliers.
|
||||
|
||||
Each subclass handles one EditOp. It validates the edit against
|
||||
the current config state, applies the mutation, and can roll back.
|
||||
"""
|
||||
|
||||
op: ClassVar[EditOp]
|
||||
|
||||
@abstractmethod
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
"""Check if the edit can be applied to the current config."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
"""Mutate the config. Must be idempotent."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
"""Restore pre-edit state. Most appliers delegate to git checkout."""
|
||||
...
|
||||
|
||||
|
||||
class EditApplierRegistry:
|
||||
"""Registry of EditApplier instances keyed by EditOp."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._appliers: dict[EditOp, EditApplier] = {}
|
||||
|
||||
def register(self, applier: EditApplier) -> None:
|
||||
"""Register an applier instance."""
|
||||
self._appliers[applier.op] = applier
|
||||
|
||||
def get(self, op: EditOp) -> EditApplier:
|
||||
"""Return the applier for the given op. Raises KeyError if not found."""
|
||||
return self._appliers[op]
|
||||
|
||||
def is_supported(self, op: EditOp) -> bool:
|
||||
"""Return True if an applier is registered for the op."""
|
||||
return op in self._appliers
|
||||
@@ -0,0 +1,194 @@
|
||||
"""Per-edit execution loop for the distillation execute phase.
|
||||
|
||||
Iterates over a plan's edits, handles tier routing, validates, and applies.
|
||||
Does NOT include the benchmark gate — that's wired in M5.
|
||||
|
||||
See spec §7.2.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
EditFewShotExemplarsApplier,
|
||||
PatchSystemPromptApplier,
|
||||
ReplaceSystemPromptApplier,
|
||||
SetAgentClassApplier,
|
||||
SetAgentParamApplier,
|
||||
)
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelForQueryClassApplier,
|
||||
SetModelParamApplier,
|
||||
)
|
||||
from openjarvis.learning.distillation.execute.appliers.lora_stub import (
|
||||
LoraStubApplier,
|
||||
)
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
AddToolToAgentApplier,
|
||||
EditToolDescriptionApplier,
|
||||
RemoveToolFromAgentApplier,
|
||||
)
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
EditApplierRegistry,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
Edit,
|
||||
EditOutcome,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _build_registry() -> EditApplierRegistry:
|
||||
"""Build and populate the default applier registry."""
|
||||
registry = EditApplierRegistry()
|
||||
registry.register(SetModelForQueryClassApplier())
|
||||
registry.register(SetModelParamApplier())
|
||||
registry.register(PatchSystemPromptApplier())
|
||||
registry.register(ReplaceSystemPromptApplier())
|
||||
registry.register(SetAgentClassApplier())
|
||||
registry.register(SetAgentParamApplier())
|
||||
registry.register(EditFewShotExemplarsApplier())
|
||||
registry.register(AddToolToAgentApplier())
|
||||
registry.register(RemoveToolFromAgentApplier())
|
||||
registry.register(EditToolDescriptionApplier())
|
||||
registry.register(LoraStubApplier())
|
||||
return registry
|
||||
|
||||
|
||||
def execute_edits(
|
||||
*,
|
||||
edits: list[Edit],
|
||||
ctx: ApplyContext,
|
||||
autonomy_mode: AutonomyMode,
|
||||
registry: EditApplierRegistry | None = None,
|
||||
) -> list[EditOutcome]:
|
||||
"""Execute a list of edits, returning outcomes for each.
|
||||
|
||||
This loop handles tier routing and validation but does NOT include
|
||||
the benchmark gate (wired in M5).
|
||||
|
||||
Parameters
|
||||
----------
|
||||
edits :
|
||||
The edits to process.
|
||||
ctx :
|
||||
Shared context with config paths.
|
||||
autonomy_mode :
|
||||
How aggressively to apply edits.
|
||||
registry :
|
||||
Optional pre-built registry. If None, builds the default.
|
||||
"""
|
||||
if registry is None:
|
||||
registry = _build_registry()
|
||||
|
||||
outcomes: list[EditOutcome] = []
|
||||
|
||||
for edit in edits:
|
||||
# Manual mode: everything goes to review
|
||||
if autonomy_mode == AutonomyMode.MANUAL:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="pending_review",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Manual tier: always skip
|
||||
if edit.risk_tier == EditRiskTier.MANUAL:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="skipped",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error="manual tier, requires explicit approval",
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Review tier in tiered mode: route to pending
|
||||
if (
|
||||
edit.risk_tier == EditRiskTier.REVIEW
|
||||
and autonomy_mode == AutonomyMode.TIERED
|
||||
):
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="pending_review",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Check if the op is supported
|
||||
if not registry.is_supported(edit.op):
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="skipped",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=f"op {edit.op.value} not implemented in v1",
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Validate
|
||||
applier = registry.get(edit.op)
|
||||
validation = applier.validate(edit, ctx)
|
||||
if not validation.ok:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=validation.reason,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Apply
|
||||
try:
|
||||
applier.apply(edit, ctx)
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="applied",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Edit %s failed: %s", edit.id, e)
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=str(e),
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
|
||||
return outcomes
|
||||
@@ -0,0 +1 @@
|
||||
"""Gate phase: benchmark-based accept/reject for edits."""
|
||||
@@ -0,0 +1,131 @@
|
||||
"""BenchmarkGate: accept or reject edits based on benchmark performance.
|
||||
|
||||
Runs the personal benchmark via a provided scorer callable, compares
|
||||
before/after snapshots, and decides accept/reject based on thresholds.
|
||||
|
||||
See spec §7.3.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable
|
||||
|
||||
from openjarvis.learning.distillation.gate.regression import regression_check
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Type for the scorer callable: takes benchmark_version, subsample_size,
|
||||
# seed and returns a BenchmarkSnapshot.
|
||||
ScorerFn = Callable[..., BenchmarkSnapshot]
|
||||
|
||||
|
||||
@dataclass
|
||||
class GateResult:
|
||||
"""Result of a benchmark gate evaluation."""
|
||||
|
||||
accepted: bool
|
||||
snapshot: BenchmarkSnapshot
|
||||
delta: float
|
||||
reason: str = ""
|
||||
|
||||
|
||||
class BenchmarkGate:
|
||||
"""Runs the personal benchmark and decides accept/reject.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
scorer :
|
||||
Callable that runs the benchmark and returns a BenchmarkSnapshot.
|
||||
Signature: ``(benchmark_version, subsample_size, seed) -> BenchmarkSnapshot``.
|
||||
benchmark_version :
|
||||
Which benchmark version to score against (locked per session).
|
||||
min_improvement :
|
||||
Minimum overall score improvement to accept (default 0.0).
|
||||
max_regression :
|
||||
Maximum per-cluster score drop before rejecting (default 0.05).
|
||||
subsample_size :
|
||||
Number of tasks to score per gate run (default 50).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
scorer: ScorerFn,
|
||||
benchmark_version: str,
|
||||
min_improvement: float = 0.0,
|
||||
max_regression: float = 0.05,
|
||||
subsample_size: int = 50,
|
||||
) -> None:
|
||||
self._scorer = scorer
|
||||
self._benchmark_version = benchmark_version
|
||||
self._min_improvement = min_improvement
|
||||
self._max_regression = max_regression
|
||||
self._subsample_size = subsample_size
|
||||
|
||||
def evaluate(
|
||||
self,
|
||||
*,
|
||||
before: BenchmarkSnapshot,
|
||||
session_seed: int,
|
||||
) -> GateResult:
|
||||
"""Run the benchmark and compare against the before snapshot.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
before :
|
||||
Snapshot captured before the edit was applied.
|
||||
session_seed :
|
||||
Deterministic seed for subsampling (same across all gate
|
||||
runs in one session).
|
||||
|
||||
Returns
|
||||
-------
|
||||
GateResult
|
||||
``accepted`` is True if the edit should be committed.
|
||||
"""
|
||||
after = self._scorer(
|
||||
benchmark_version=self._benchmark_version,
|
||||
subsample_size=self._subsample_size,
|
||||
seed=session_seed,
|
||||
)
|
||||
|
||||
delta = after.overall_score - before.overall_score
|
||||
|
||||
# Check regression
|
||||
reg = regression_check(before, after, max_regression=self._max_regression)
|
||||
if reg.has_regression:
|
||||
clusters_str = ", ".join(
|
||||
f"{cid}: {d:+.3f}" for cid, d in reg.regressed_clusters.items()
|
||||
)
|
||||
reason = f"regression in clusters: {clusters_str}"
|
||||
logger.info("Gate rejected: %s", reason)
|
||||
return GateResult(
|
||||
accepted=False,
|
||||
snapshot=after,
|
||||
delta=delta,
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
# Check improvement
|
||||
if delta <= self._min_improvement:
|
||||
reason = (
|
||||
f"no improvement: delta={delta:.4f}, "
|
||||
f"min_improvement={self._min_improvement}"
|
||||
)
|
||||
logger.info("Gate rejected: %s", reason)
|
||||
return GateResult(
|
||||
accepted=False,
|
||||
snapshot=after,
|
||||
delta=delta,
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
logger.info("Gate accepted: delta=%.4f", delta)
|
||||
return GateResult(
|
||||
accepted=True,
|
||||
snapshot=after,
|
||||
delta=delta,
|
||||
)
|
||||
@@ -0,0 +1,109 @@
|
||||
"""Cold start detection and bootstrap for the distillation subsystem.
|
||||
|
||||
Day one: no traces, no benchmark. The system must not crash and must
|
||||
give the user a clear message about what's needed. This module provides
|
||||
readiness checks and the bootstrap logic.
|
||||
|
||||
See spec §13.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReadinessResult:
|
||||
"""Result of a readiness check."""
|
||||
|
||||
ready: bool
|
||||
message: str
|
||||
trace_count: int = 0
|
||||
high_feedback_count: int = 0
|
||||
|
||||
|
||||
def check_readiness(
|
||||
trace_store: Any,
|
||||
min_traces: int = 20,
|
||||
) -> ReadinessResult:
|
||||
"""Check if there are enough traces to run a learning session.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
trace_store :
|
||||
TraceStore instance.
|
||||
min_traces :
|
||||
Minimum total trace count required.
|
||||
|
||||
Returns
|
||||
-------
|
||||
ReadinessResult
|
||||
``ready`` is True if there are enough traces.
|
||||
"""
|
||||
count = trace_store.count()
|
||||
if count < min_traces:
|
||||
return ReadinessResult(
|
||||
ready=False,
|
||||
message=(
|
||||
f"Not enough traces yet to learn from. "
|
||||
f"Have {count}, need at least {min_traces}. "
|
||||
f"Use OpenJarvis for a while and try again."
|
||||
),
|
||||
trace_count=count,
|
||||
)
|
||||
return ReadinessResult(
|
||||
ready=True,
|
||||
message=f"Ready: {count} traces available.",
|
||||
trace_count=count,
|
||||
)
|
||||
|
||||
|
||||
def check_benchmark_ready(
|
||||
trace_store: Any,
|
||||
min_feedback: float = 0.7,
|
||||
min_samples: int = 10,
|
||||
) -> ReadinessResult:
|
||||
"""Check if there are enough high-feedback traces for a benchmark.
|
||||
|
||||
The bootstrap benchmark needs at least ``min_samples`` traces with
|
||||
feedback >= ``min_feedback``.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
trace_store :
|
||||
TraceStore instance.
|
||||
min_feedback :
|
||||
Minimum feedback score for benchmark-quality traces.
|
||||
min_samples :
|
||||
Minimum number of high-feedback traces needed.
|
||||
"""
|
||||
# Query traces with high feedback
|
||||
traces = trace_store.list_traces(limit=min_samples * 2)
|
||||
high_feedback = [
|
||||
t
|
||||
for t in traces
|
||||
if getattr(t, "feedback", None) is not None and t.feedback >= min_feedback
|
||||
]
|
||||
count = len(high_feedback)
|
||||
|
||||
if count < min_samples:
|
||||
return ReadinessResult(
|
||||
ready=False,
|
||||
message=(
|
||||
f"Personal benchmark needs at least {min_samples} "
|
||||
f"high-feedback traces (feedback >= {min_feedback}). "
|
||||
f"Have {count} so far. Will be populated automatically."
|
||||
),
|
||||
trace_count=trace_store.count(),
|
||||
high_feedback_count=count,
|
||||
)
|
||||
return ReadinessResult(
|
||||
ready=True,
|
||||
message=f"Benchmark ready: {count} high-feedback traces available.",
|
||||
trace_count=trace_store.count(),
|
||||
high_feedback_count=count,
|
||||
)
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Per-cluster regression detection for the benchmark gate.
|
||||
|
||||
After an edit is applied, the gate compares the before and after
|
||||
BenchmarkSnapshots. If any cluster's score dropped by more than
|
||||
``max_regression``, the edit is rejected.
|
||||
|
||||
See spec §7.3.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
|
||||
@dataclass
|
||||
class RegressionResult:
|
||||
"""Result of a regression check."""
|
||||
|
||||
has_regression: bool
|
||||
regressed_clusters: dict[str, float] = field(default_factory=dict)
|
||||
"""Cluster id → negative delta for clusters that regressed beyond threshold."""
|
||||
|
||||
|
||||
def regression_check(
|
||||
before: BenchmarkSnapshot,
|
||||
after: BenchmarkSnapshot,
|
||||
max_regression: float = 0.05,
|
||||
) -> RegressionResult:
|
||||
"""Check if any cluster regressed beyond the threshold.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
before :
|
||||
Benchmark snapshot before the edit.
|
||||
after :
|
||||
Benchmark snapshot after the edit.
|
||||
max_regression :
|
||||
Maximum allowed per-cluster score drop (default 0.05).
|
||||
|
||||
Returns
|
||||
-------
|
||||
RegressionResult
|
||||
``has_regression`` is True if any cluster dropped more than
|
||||
``max_regression``. ``regressed_clusters`` maps cluster ids
|
||||
to their negative deltas.
|
||||
"""
|
||||
regressed: dict[str, float] = {}
|
||||
|
||||
for cluster_id, before_score in before.cluster_scores.items():
|
||||
after_score = after.cluster_scores.get(cluster_id, 0.0)
|
||||
delta = after_score - before_score
|
||||
if delta < -max_regression:
|
||||
regressed[cluster_id] = delta
|
||||
|
||||
return RegressionResult(
|
||||
has_regression=bool(regressed),
|
||||
regressed_clusters=regressed,
|
||||
)
|
||||
@@ -0,0 +1,372 @@
|
||||
"""Pydantic models and enums for the distillation subsystem.
|
||||
|
||||
This module defines the typed vocabulary used by the diagnose, plan, execute,
|
||||
and record phases. Three model families:
|
||||
|
||||
- Enums: pillar / risk tier / op / trigger kind / autonomy mode / session status
|
||||
- Edit + LearningPlan + FailureCluster: the teacher's frozen output
|
||||
- LearningSession + EditOutcome + BenchmarkSnapshot: the durable session record
|
||||
|
||||
See spec §4 for the data model rationale.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enums
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class EditPillar(str, Enum):
|
||||
"""Which OpenJarvis pillar an edit targets."""
|
||||
|
||||
INTELLIGENCE = "intelligence"
|
||||
AGENT = "agent"
|
||||
TOOLS = "tools"
|
||||
ENGINE = "engine" # designed-for; no v1 appliers ship
|
||||
|
||||
|
||||
class EditRiskTier(str, Enum):
|
||||
"""How an edit gets applied: auto, review queue, or manual-only."""
|
||||
|
||||
AUTO = "auto"
|
||||
REVIEW = "review"
|
||||
MANUAL = "manual"
|
||||
|
||||
|
||||
class EditOp(str, Enum):
|
||||
"""The set of typed operations a teacher can propose.
|
||||
|
||||
Each op corresponds to one EditApplier in v1 (or a refusing stub for
|
||||
deferred ops). The teacher cannot invent new ops — only choose from this
|
||||
set.
|
||||
"""
|
||||
|
||||
# Intelligence
|
||||
SET_MODEL_FOR_QUERY_CLASS = "set_model_for_query_class"
|
||||
SET_MODEL_PARAM = "set_model_param"
|
||||
|
||||
# Agent
|
||||
PATCH_SYSTEM_PROMPT = "patch_system_prompt"
|
||||
REPLACE_SYSTEM_PROMPT = "replace_system_prompt"
|
||||
SET_AGENT_CLASS = "set_agent_class"
|
||||
SET_AGENT_PARAM = "set_agent_param"
|
||||
EDIT_FEW_SHOT_EXEMPLARS = "edit_few_shot_exemplars"
|
||||
|
||||
# Tools
|
||||
ADD_TOOL_TO_AGENT = "add_tool_to_agent"
|
||||
REMOVE_TOOL_FROM_AGENT = "remove_tool_from_agent"
|
||||
EDIT_TOOL_DESCRIPTION = "edit_tool_description"
|
||||
|
||||
# v2 placeholder — planner can emit, executor refuses with NotImplementedError
|
||||
LORA_FINETUNE = "lora_finetune"
|
||||
|
||||
|
||||
class TriggerKind(str, Enum):
|
||||
"""What kicked off a learning session."""
|
||||
|
||||
SCHEDULED = "scheduled"
|
||||
CLUSTER = "cluster"
|
||||
USER_FLAG = "user_flag"
|
||||
ON_DEMAND = "on_demand"
|
||||
|
||||
|
||||
class AutonomyMode(str, Enum):
|
||||
"""How aggressively the orchestrator applies edits without review."""
|
||||
|
||||
AUTO = "auto" # all tiers auto-apply, ignore review tier
|
||||
TIERED = "tiered" # default: respect per-edit risk tier
|
||||
MANUAL = "manual" # everything goes to review queue (dry-run mode)
|
||||
|
||||
|
||||
class SessionStatus(str, Enum):
|
||||
"""Lifecycle states for a LearningSession.
|
||||
|
||||
See spec §7.7 for the transition rules.
|
||||
"""
|
||||
|
||||
INITIATED = "initiated"
|
||||
DIAGNOSING = "diagnosing"
|
||||
PLANNING = "planning"
|
||||
EXECUTING = "executing"
|
||||
AWAITING_REVIEW = "awaiting_review"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
ROLLED_BACK = "rolled_back"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Edit — atomic unit of change
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class Edit(BaseModel):
|
||||
"""One atomic edit to the OpenJarvis harness.
|
||||
|
||||
Emitted by the LearningPlanner, consumed by an EditApplier. The teacher
|
||||
proposes the op, target, payload, rationale, and references; the planner
|
||||
overwrites ``risk_tier`` deterministically from the (pillar, op) lookup
|
||||
table — the teacher cannot pick its own tier.
|
||||
|
||||
See spec §4.1.
|
||||
"""
|
||||
|
||||
id: str = Field(
|
||||
...,
|
||||
description="UUID for this edit; also used as a footer in git commits.",
|
||||
)
|
||||
pillar: EditPillar = Field(..., description="Which OpenJarvis pillar is touched.")
|
||||
op: EditOp = Field(..., description="The typed operation to perform.")
|
||||
target: str = Field(
|
||||
...,
|
||||
description="Dotted path to the target, e.g. 'agents.simple.system_prompt'.",
|
||||
)
|
||||
payload: dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description="Op-specific arguments. Schema depends on op.",
|
||||
)
|
||||
rationale: str = Field(
|
||||
...,
|
||||
description="Teacher's natural-language reason for this edit.",
|
||||
)
|
||||
expected_improvement: str = Field(
|
||||
...,
|
||||
description="Which failure cluster id this edit is intended to address.",
|
||||
)
|
||||
risk_tier: EditRiskTier = Field(
|
||||
...,
|
||||
description="Set by the planner from a (pillar, op) lookup table.",
|
||||
)
|
||||
references: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Trace ids or benchmark task ids that justify this edit.",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FailureCluster — a group of related failures with a shared root cause
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class FailureCluster(BaseModel):
|
||||
"""A group of failing traces that share a hypothesised root cause.
|
||||
|
||||
Populated by the teacher in the diagnose phase. The student/teacher rates
|
||||
must come from real ``run_student_on_task`` and ``run_self_on_task`` calls
|
||||
against benchmark tasks (see spec §5.3); clusters where both rates are
|
||||
missing or zero are dropped by the planner.
|
||||
|
||||
See spec §4.2.
|
||||
"""
|
||||
|
||||
id: str
|
||||
description: str = Field(..., description="Short human description of the cluster.")
|
||||
sample_trace_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Trace ids that exemplify this cluster (>= 3 typical).",
|
||||
)
|
||||
student_failure_rate: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Local student's failure rate on benchmark tasks in this cluster.",
|
||||
)
|
||||
teacher_success_rate: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Frontier teacher's success rate on the same tasks.",
|
||||
)
|
||||
skill_gap: str = Field(
|
||||
...,
|
||||
description="Teacher's qualitative analysis of what the student is missing.",
|
||||
)
|
||||
addressed_by_edit_ids: list[str] = Field(
|
||||
default_factory=list,
|
||||
description="Ids of edits in the LearningPlan that target this cluster.",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LearningPlan — frozen output of the planning phase
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class LearningPlan(BaseModel):
|
||||
"""The teacher's frozen plan of edits for a learning session.
|
||||
|
||||
Once written to ``<session_dir>/plan.json`` this is immutable. The
|
||||
execution layer reads this file and does not re-prompt the teacher.
|
||||
|
||||
See spec §4.2 and §6.
|
||||
"""
|
||||
|
||||
session_id: str = Field(..., description="Owning session id.")
|
||||
diagnosis_summary: str = Field(
|
||||
...,
|
||||
description="Markdown narrative analysis from the teacher (~500-2000 words).",
|
||||
)
|
||||
failure_clusters: list[FailureCluster] = Field(
|
||||
default_factory=list,
|
||||
description="Clusters identified in phase 1.",
|
||||
)
|
||||
edits: list[Edit] = Field(
|
||||
default_factory=list,
|
||||
description="Typed edit list emitted by the planner.",
|
||||
)
|
||||
teacher_model: str = Field(
|
||||
...,
|
||||
description="Frontier model id used as the teacher (e.g. 'claude-opus-4-6').",
|
||||
)
|
||||
estimated_cost_usd: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
description="Total teacher API cost estimate for this session.",
|
||||
)
|
||||
created_at: datetime = Field(
|
||||
...,
|
||||
description="When the plan was finalized.",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BenchmarkSnapshot — one personal-benchmark run result
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BenchmarkSnapshot(BaseModel):
|
||||
"""A point-in-time score from running the personal benchmark.
|
||||
|
||||
Two of these live on every LearningSession: one captured before any edits
|
||||
apply, one captured after. The version is locked at session start so the
|
||||
delta is interpretable even if the benchmark is refreshed mid-session.
|
||||
|
||||
See spec §4.3 and §9.
|
||||
"""
|
||||
|
||||
benchmark_version: str = Field(
|
||||
...,
|
||||
description="Personal benchmark version (e.g. 'personal_v3').",
|
||||
)
|
||||
overall_score: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
description="Mean per-task score across the benchmark.",
|
||||
)
|
||||
cluster_scores: dict[str, float] = Field(
|
||||
default_factory=dict,
|
||||
description="Mean score per failure cluster.",
|
||||
)
|
||||
task_count: int = Field(
|
||||
...,
|
||||
ge=0,
|
||||
description="Number of tasks scored in this snapshot.",
|
||||
)
|
||||
elapsed_seconds: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
description="Wall-clock time of the benchmark run.",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EditOutcome — what happened when one edit was processed
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
EditOutcomeStatus = Literal[
|
||||
"applied",
|
||||
"rejected_by_gate",
|
||||
"pending_review",
|
||||
"rejected_by_user",
|
||||
"rolled_back",
|
||||
"skipped",
|
||||
]
|
||||
|
||||
|
||||
class EditOutcome(BaseModel):
|
||||
"""Result of attempting to apply one Edit.
|
||||
|
||||
Persisted both in the SessionStore SQLite table and as part of the
|
||||
session.json artifact. The status literal is the canonical lifecycle for
|
||||
each edit.
|
||||
|
||||
See spec §4.3.
|
||||
"""
|
||||
|
||||
edit_id: str
|
||||
status: EditOutcomeStatus
|
||||
benchmark_delta: float | None = Field(
|
||||
default=None,
|
||||
description="Overall score change from this edit. None if not gated.",
|
||||
)
|
||||
cluster_deltas: dict[str, float] = Field(
|
||||
default_factory=dict,
|
||||
description="Per-cluster score change from this edit.",
|
||||
)
|
||||
error: str | None = Field(
|
||||
default=None,
|
||||
description="Error message if the edit was rejected or failed.",
|
||||
)
|
||||
applied_at: datetime | None = Field(
|
||||
default=None,
|
||||
description="When the edit was committed to the checkpoint repo.",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LearningSession — durable record of one full diagnose→record loop
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class LearningSession(BaseModel):
|
||||
"""The durable record of one distillation session.
|
||||
|
||||
Persisted in two places: `<session_dir>/session.json` (authoritative) and
|
||||
the SQLite SessionStore (queryable index). When in doubt, prefer the JSON
|
||||
file — SQLite can be rebuilt from the JSON files.
|
||||
|
||||
See spec §4.3, §7.7 (status transitions), §8.
|
||||
"""
|
||||
|
||||
id: str
|
||||
parent_session_id: str | None = Field(
|
||||
default=None,
|
||||
description="Id of the session this one is a follow-up to, if any.",
|
||||
)
|
||||
trigger: TriggerKind
|
||||
trigger_metadata: dict[str, Any] = Field(default_factory=dict)
|
||||
status: SessionStatus
|
||||
autonomy_mode: AutonomyMode
|
||||
started_at: datetime
|
||||
ended_at: datetime | None = None
|
||||
diagnosis_path: Path
|
||||
plan_path: Path
|
||||
benchmark_before: BenchmarkSnapshot
|
||||
benchmark_after: BenchmarkSnapshot | None = None
|
||||
edit_outcomes: list[EditOutcome] = Field(default_factory=list)
|
||||
git_checkpoint_pre: str = Field(
|
||||
...,
|
||||
description="Commit sha at session start (baseline commit).",
|
||||
)
|
||||
git_checkpoint_post: str | None = Field(
|
||||
default=None,
|
||||
description="Commit sha after edits applied; None until executing finishes.",
|
||||
)
|
||||
teacher_cost_usd: float = Field(
|
||||
...,
|
||||
ge=0.0,
|
||||
description="Accumulated teacher API spend for this session.",
|
||||
)
|
||||
error: str | None = Field(
|
||||
default=None,
|
||||
description="If status is FAILED, the error message that caused it.",
|
||||
)
|
||||
@@ -0,0 +1,418 @@
|
||||
"""DistillationOrchestrator: top-level driver for a learning session.
|
||||
|
||||
Wires diagnose (M2) → plan (M3) → execute (M4) → gate (M5) into a
|
||||
single ``run(trigger)`` method. All dependencies are injected.
|
||||
|
||||
See spec §3, §7.2, §7.7.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from openjarvis.learning.distillation.diagnose.runner import DiagnosisRunner
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.execute.loop import _build_registry
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import BenchmarkGate
|
||||
from openjarvis.learning.distillation.gate.cold_start import check_readiness
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
BenchmarkSnapshot,
|
||||
EditOutcome,
|
||||
EditRiskTier,
|
||||
LearningSession,
|
||||
SessionStatus,
|
||||
)
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DistillationOrchestrator:
|
||||
"""Top-level driver for a distillation learning session.
|
||||
|
||||
All dependencies are injected so tests can mock everything.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
teacher_engine: Any,
|
||||
teacher_model: str,
|
||||
trace_store: Any,
|
||||
benchmark_samples: list,
|
||||
student_runner: Any,
|
||||
judge: Any,
|
||||
session_store: Any,
|
||||
checkpoint_store: Any,
|
||||
openjarvis_home: Path,
|
||||
autonomy_mode: AutonomyMode = AutonomyMode.TIERED,
|
||||
scorer: Callable[..., BenchmarkSnapshot] | None = None,
|
||||
benchmark_version: str = "personal_v1",
|
||||
min_traces: int = 20,
|
||||
max_cost_usd: float = 5.0,
|
||||
max_tool_calls: int = 30,
|
||||
min_improvement: float = 0.0,
|
||||
max_regression: float = 0.05,
|
||||
subsample_size: int = 50,
|
||||
) -> None:
|
||||
self._engine = teacher_engine
|
||||
self._model = teacher_model
|
||||
self._trace_store = trace_store
|
||||
self._benchmark_samples = benchmark_samples
|
||||
self._student_runner = student_runner
|
||||
self._judge = judge
|
||||
self._session_store = session_store
|
||||
self._checkpoint_store = checkpoint_store
|
||||
self._home = Path(openjarvis_home)
|
||||
self._autonomy = autonomy_mode
|
||||
self._scorer = scorer
|
||||
self._bench_version = benchmark_version
|
||||
self._min_traces = min_traces
|
||||
self._max_cost = max_cost_usd
|
||||
self._max_tool_calls = max_tool_calls
|
||||
self._min_improvement = min_improvement
|
||||
self._max_regression = max_regression
|
||||
self._subsample_size = subsample_size
|
||||
|
||||
def run(self, trigger: Any) -> LearningSession:
|
||||
"""Execute a full distillation session.
|
||||
|
||||
Returns the completed LearningSession.
|
||||
"""
|
||||
ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%S")
|
||||
session_id = f"session-{ts}_{uuid.uuid4().hex[:8]}"
|
||||
session_dir = self._home / "learning" / "sessions" / session_id
|
||||
session_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Initialize session
|
||||
pre_sha = self._checkpoint_store.current_sha()
|
||||
session = LearningSession(
|
||||
id=session_id,
|
||||
trigger=trigger.kind,
|
||||
trigger_metadata=trigger.metadata,
|
||||
status=SessionStatus.INITIATED,
|
||||
autonomy_mode=self._autonomy,
|
||||
started_at=datetime.now(timezone.utc),
|
||||
diagnosis_path=session_dir / "diagnosis.md",
|
||||
plan_path=session_dir / "plan.json",
|
||||
benchmark_before=BenchmarkSnapshot(
|
||||
benchmark_version=self._bench_version,
|
||||
overall_score=0.0,
|
||||
cluster_scores={},
|
||||
task_count=0,
|
||||
elapsed_seconds=0.0,
|
||||
),
|
||||
git_checkpoint_pre=pre_sha,
|
||||
teacher_cost_usd=0.0,
|
||||
)
|
||||
|
||||
try:
|
||||
# Cold start check
|
||||
readiness = check_readiness(self._trace_store, min_traces=self._min_traces)
|
||||
if not readiness.ready:
|
||||
session = session.model_copy(
|
||||
update={
|
||||
"status": SessionStatus.FAILED,
|
||||
"error": readiness.message,
|
||||
"ended_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
self._session_store.save_session(session)
|
||||
return session
|
||||
|
||||
# Capture benchmark before
|
||||
if self._scorer is not None:
|
||||
before_snap = self._scorer(
|
||||
benchmark_version=self._bench_version,
|
||||
subsample_size=self._subsample_size,
|
||||
seed=hash(session_id) % (2**31),
|
||||
)
|
||||
session = session.model_copy(update={"benchmark_before": before_snap})
|
||||
|
||||
# Phase 1: Diagnose
|
||||
session = session.model_copy(update={"status": SessionStatus.DIAGNOSING})
|
||||
self._session_store.save_session(session)
|
||||
|
||||
diagnosis_runner = DiagnosisRunner(
|
||||
teacher_engine=self._engine,
|
||||
teacher_model=self._model,
|
||||
trace_store=self._trace_store,
|
||||
benchmark_samples=self._benchmark_samples,
|
||||
student_runner=self._student_runner,
|
||||
judge=self._judge,
|
||||
session_dir=session_dir,
|
||||
session_id=session_id,
|
||||
config={
|
||||
"config_path": self._home / "config.toml",
|
||||
"openjarvis_home": self._home,
|
||||
},
|
||||
max_turns=self._max_tool_calls,
|
||||
max_cost_usd=self._max_cost,
|
||||
)
|
||||
diag_result = diagnosis_runner.run()
|
||||
cost = diag_result.cost_usd
|
||||
|
||||
if not diag_result.clusters:
|
||||
session = session.model_copy(
|
||||
update={
|
||||
"status": SessionStatus.FAILED,
|
||||
"error": "diagnosis produced no actionable clusters",
|
||||
"teacher_cost_usd": cost,
|
||||
"ended_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
self._session_store.save_session(session)
|
||||
return session
|
||||
|
||||
# Phase 2: Plan
|
||||
session = session.model_copy(update={"status": SessionStatus.PLANNING})
|
||||
self._session_store.save_session(session)
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=self._engine,
|
||||
teacher_model=self._model,
|
||||
session_id=session_id,
|
||||
session_dir=session_dir,
|
||||
prompt_reader=lambda t: self._read_prompt(t),
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md=diag_result.diagnosis_md,
|
||||
clusters=diag_result.clusters,
|
||||
)
|
||||
cost += plan.estimated_cost_usd
|
||||
|
||||
# Phase 3: Execute
|
||||
session = session.model_copy(update={"status": SessionStatus.EXECUTING})
|
||||
self._session_store.save_session(session)
|
||||
|
||||
ctx = ApplyContext(
|
||||
openjarvis_home=self._home,
|
||||
session_id=session_id,
|
||||
)
|
||||
registry = _build_registry()
|
||||
|
||||
gate: BenchmarkGate | None = None
|
||||
if self._scorer is not None:
|
||||
gate = BenchmarkGate(
|
||||
scorer=self._scorer,
|
||||
benchmark_version=self._bench_version,
|
||||
min_improvement=self._min_improvement,
|
||||
max_regression=self._max_regression,
|
||||
subsample_size=self._subsample_size,
|
||||
)
|
||||
|
||||
session_seed = hash(session_id) % (2**31)
|
||||
outcomes: list[EditOutcome] = []
|
||||
|
||||
for edit in plan.edits:
|
||||
# Manual autonomy mode: everything goes to review
|
||||
if self._autonomy == AutonomyMode.MANUAL:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="pending_review",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Manual risk tier: always skip
|
||||
if edit.risk_tier == EditRiskTier.MANUAL:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="skipped",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error="manual tier, requires explicit approval",
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Review tier in tiered mode: route to pending
|
||||
if (
|
||||
edit.risk_tier == EditRiskTier.REVIEW
|
||||
and self._autonomy == AutonomyMode.TIERED
|
||||
):
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="pending_review",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Check if the op is supported
|
||||
if not registry.is_supported(edit.op):
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="skipped",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=f"op {edit.op.value} not implemented in v1",
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Validate
|
||||
applier = registry.get(edit.op)
|
||||
validation = applier.validate(edit, ctx)
|
||||
if not validation.ok:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=validation.reason,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Apply the edit
|
||||
try:
|
||||
applier.apply(edit, ctx)
|
||||
except Exception as exc:
|
||||
logger.warning("Edit %s apply failed: %s", edit.id, exc)
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=str(exc),
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# If no scorer, accept directly (backward-compat with tests)
|
||||
if gate is None:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="applied",
|
||||
benchmark_delta=None,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# Run the benchmark gate
|
||||
before_snap = session.benchmark_before
|
||||
gate_result = gate.evaluate(
|
||||
before=before_snap,
|
||||
session_seed=session_seed,
|
||||
)
|
||||
|
||||
if gate_result.accepted:
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="applied",
|
||||
benchmark_delta=gate_result.delta,
|
||||
cluster_deltas={},
|
||||
error=None,
|
||||
applied_at=datetime.now(timezone.utc),
|
||||
)
|
||||
)
|
||||
else:
|
||||
# Gate rejected — rollback the edit
|
||||
try:
|
||||
applier.rollback(edit, ctx)
|
||||
except Exception as rb_exc:
|
||||
logger.warning(
|
||||
"Edit %s rollback failed: %s", edit.id, rb_exc
|
||||
)
|
||||
outcomes.append(
|
||||
EditOutcome(
|
||||
edit_id=edit.id,
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=gate_result.delta,
|
||||
cluster_deltas={},
|
||||
error=gate_result.reason,
|
||||
applied_at=None,
|
||||
)
|
||||
)
|
||||
|
||||
# Enqueue pending_review edits
|
||||
pending_queue = PendingQueue(self._home / "learning" / "pending_review")
|
||||
has_pending = False
|
||||
for outcome, edit in zip(outcomes, plan.edits):
|
||||
if outcome.status == "pending_review":
|
||||
pending_queue.enqueue(session_id, edit)
|
||||
has_pending = True
|
||||
|
||||
# Capture benchmark after
|
||||
after_snap = None
|
||||
if self._scorer is not None:
|
||||
after_snap = self._scorer(
|
||||
benchmark_version=self._bench_version,
|
||||
subsample_size=self._subsample_size,
|
||||
seed=session_seed,
|
||||
)
|
||||
|
||||
# Determine final status
|
||||
if has_pending:
|
||||
final_status = SessionStatus.AWAITING_REVIEW
|
||||
else:
|
||||
final_status = SessionStatus.COMPLETED
|
||||
|
||||
post_sha = self._checkpoint_store.current_sha()
|
||||
session = session.model_copy(
|
||||
update={
|
||||
"status": final_status,
|
||||
"edit_outcomes": outcomes,
|
||||
"benchmark_after": after_snap,
|
||||
"git_checkpoint_post": post_sha,
|
||||
"teacher_cost_usd": cost,
|
||||
"ended_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("Session %s failed: %s", session_id, e)
|
||||
session = session.model_copy(
|
||||
update={
|
||||
"status": SessionStatus.FAILED,
|
||||
"error": str(e),
|
||||
"ended_at": datetime.now(timezone.utc),
|
||||
}
|
||||
)
|
||||
|
||||
self._session_store.save_session(session)
|
||||
|
||||
# Write session.json artifact
|
||||
artifact_path = session_dir / "session.json"
|
||||
artifact_path.write_text(session.model_dump_json(indent=2), encoding="utf-8")
|
||||
|
||||
return session
|
||||
|
||||
def _read_prompt(self, target: str) -> str:
|
||||
"""Read a prompt file from the config tree."""
|
||||
parts = target.split(".")
|
||||
if len(parts) >= 2:
|
||||
agent_name = parts[1]
|
||||
path = self._home / "agents" / agent_name / "system_prompt.md"
|
||||
if path.exists():
|
||||
return path.read_text(encoding="utf-8")
|
||||
return ""
|
||||
@@ -0,0 +1,74 @@
|
||||
"""Pending review queue for edits awaiting user approval.
|
||||
|
||||
Edits in the ``review`` tier (when autonomy mode is ``tiered``) are
|
||||
written here as JSON files. The user reviews them via ``jarvis learning
|
||||
review`` and approves or rejects.
|
||||
|
||||
See spec §7.5.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PendingQueue:
|
||||
"""File-based queue for pending review edits.
|
||||
|
||||
Each edit is stored as ``<queue_dir>/<session_id>__<edit_id>.json``.
|
||||
"""
|
||||
|
||||
def __init__(self, queue_dir: Path) -> None:
|
||||
self._dir = Path(queue_dir)
|
||||
self._dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def enqueue(self, session_id: str, edit: Edit) -> Path:
|
||||
"""Write an edit to the pending queue. Returns the file path."""
|
||||
filename = f"{session_id}__{edit.id}.json"
|
||||
path = self._dir / filename
|
||||
data = {
|
||||
"session_id": session_id,
|
||||
"edit": json.loads(edit.model_dump_json()),
|
||||
}
|
||||
path.write_text(json.dumps(data, indent=2), encoding="utf-8")
|
||||
logger.info("Enqueued edit %s for review", edit.id)
|
||||
return path
|
||||
|
||||
def list_pending(self) -> list[dict[str, Any]]:
|
||||
"""Return all pending edits as dicts."""
|
||||
results = []
|
||||
for path in sorted(self._dir.glob("*.json")):
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
results.append(data)
|
||||
except (json.JSONDecodeError, OSError):
|
||||
logger.warning("Skipping corrupt pending file: %s", path)
|
||||
return results
|
||||
|
||||
def get(self, session_id: str, edit_id: str) -> dict[str, Any] | None:
|
||||
"""Return a specific pending edit, or None."""
|
||||
filename = f"{session_id}__{edit_id}.json"
|
||||
path = self._dir / filename
|
||||
if not path.exists():
|
||||
return None
|
||||
try:
|
||||
return json.loads(path.read_text(encoding="utf-8"))
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return None
|
||||
|
||||
def resolve(self, session_id: str, edit_id: str) -> bool:
|
||||
"""Remove a pending edit (approved or rejected). Returns True if found."""
|
||||
filename = f"{session_id}__{edit_id}.json"
|
||||
path = self._dir / filename
|
||||
if path.exists():
|
||||
path.unlink()
|
||||
logger.info("Resolved pending edit %s", edit_id)
|
||||
return True
|
||||
return False
|
||||
@@ -0,0 +1 @@
|
||||
"""Plan phase: converts diagnosis into a frozen LearningPlan."""
|
||||
@@ -0,0 +1,281 @@
|
||||
"""LearningPlanner: converts a diagnosis into a frozen LearningPlan.
|
||||
|
||||
Makes a single structured-output teacher call (no tools, no multi-turn)
|
||||
to generate typed edits for each failure cluster. Post-processes the
|
||||
edits with risk tier assignment and patch/replace downgrade.
|
||||
|
||||
See spec §6.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
FailureCluster,
|
||||
LearningPlan,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
maybe_downgrade_to_replace,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tiers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_PLANNER_SYSTEM_PROMPT = """\
|
||||
You are a meta-engineer planning improvements to a local AI assistant called \
|
||||
OpenJarvis. You have been given a diagnosis of the student's failure patterns.
|
||||
|
||||
Your job: for each surviving failure cluster, propose 1-3 edits from the \
|
||||
available operation set that would address the cluster's skill gap.
|
||||
|
||||
IMPORTANT: Each edit's payload must EXACTLY match the schema below. \
|
||||
Edits with missing required payload keys will be rejected.
|
||||
|
||||
Available operations with their EXACT payload schemas:
|
||||
|
||||
INTELLIGENCE pillar:
|
||||
- set_model_for_query_class: {{"query_class": "math", "model": "qwen3.5:27b"}}
|
||||
- set_model_param: {{"model": "qwen3.5:9b", "param": "temperature", "value": 0.3}}
|
||||
|
||||
AGENT pillar:
|
||||
- replace_system_prompt: {{"new_content": "You are a helpful assistant.\\n..."}}
|
||||
- patch_system_prompt: {{"diff": "--- a/prompt.md\\n+++ b/prompt.md\\n@@ ...\\n"}}
|
||||
- set_agent_class: {{"agent": "simple", "new_class": "react"}}
|
||||
- set_agent_param: {{"agent": "native_react", "param": "max_turns", "value": 10}}
|
||||
- edit_few_shot_exemplars: {{"agent": "native_react", \
|
||||
"exemplars": [{{"input": "Q", "output": "A"}}]}}
|
||||
|
||||
TOOLS pillar:
|
||||
- add_tool_to_agent: {{"agent": "native_react", "tool_name": "calculator"}}
|
||||
- remove_tool_from_agent: {{"agent": "native_react", "tool_name": "shell_exec"}}
|
||||
- edit_tool_description: {{"tool_name": "web_search", \
|
||||
"new_description": "Search the web for..."}}
|
||||
|
||||
Each edit object must have ALL of these fields:
|
||||
- id (string, e.g. "edit_001")
|
||||
- pillar ("intelligence", "agent", or "tools")
|
||||
- op (one of the operation names above)
|
||||
- target (dotted path, e.g. "agents.native_react.system_prompt")
|
||||
- payload (object matching the schema above for the chosen op)
|
||||
- rationale (string explaining why)
|
||||
- expected_improvement (cluster id this addresses)
|
||||
- risk_tier ("auto" for safe changes, "review" for prompts)
|
||||
- references (list of trace ids that justify this edit)
|
||||
|
||||
Respond with ONLY a JSON object: {{"edits": [...]}}
|
||||
"""
|
||||
|
||||
|
||||
def _validate_clusters(
|
||||
clusters: list[FailureCluster],
|
||||
) -> tuple[list[FailureCluster], list[FailureCluster]]:
|
||||
"""Split clusters into surviving and dropped.
|
||||
|
||||
Drops clusters where both student_failure_rate and teacher_success_rate
|
||||
are 0 (no evidence). Dropped clusters get a marker in skill_gap.
|
||||
"""
|
||||
surviving = []
|
||||
dropped = []
|
||||
for cluster in clusters:
|
||||
if cluster.student_failure_rate == 0.0 and cluster.teacher_success_rate == 0.0:
|
||||
marked = cluster.model_copy(
|
||||
update={
|
||||
"skill_gap": (
|
||||
"dropped: insufficient evidence"
|
||||
f" (original: {cluster.skill_gap})"
|
||||
),
|
||||
"addressed_by_edit_ids": [],
|
||||
}
|
||||
)
|
||||
dropped.append(marked)
|
||||
else:
|
||||
surviving.append(cluster)
|
||||
return surviving, dropped
|
||||
|
||||
|
||||
class LearningPlanner:
|
||||
"""Converts a diagnosis into a frozen LearningPlan.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
teacher_engine :
|
||||
CloudEngine (or mock) for the planner call.
|
||||
teacher_model :
|
||||
Frontier model id.
|
||||
session_id :
|
||||
Current session id.
|
||||
session_dir :
|
||||
Path for persisting plan.json and teacher_traces/plan.jsonl.
|
||||
prompt_reader :
|
||||
Callable that takes a target string and returns the current prompt
|
||||
content. Used by the patch/replace downgrade logic.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
teacher_engine: Any,
|
||||
teacher_model: str,
|
||||
session_id: str,
|
||||
session_dir: Path,
|
||||
prompt_reader: Callable[[str], str],
|
||||
) -> None:
|
||||
self._engine = teacher_engine
|
||||
self._model = teacher_model
|
||||
self._session_id = session_id
|
||||
self._session_dir = Path(session_dir)
|
||||
self._prompt_reader = prompt_reader
|
||||
|
||||
def run(
|
||||
self,
|
||||
*,
|
||||
diagnosis_md: str,
|
||||
clusters: list[FailureCluster],
|
||||
) -> LearningPlan:
|
||||
"""Execute the plan phase.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
diagnosis_md :
|
||||
The teacher's diagnosis markdown from phase 1.
|
||||
clusters :
|
||||
Failure clusters from phase 1.
|
||||
|
||||
Returns
|
||||
-------
|
||||
LearningPlan
|
||||
The frozen plan, also persisted to ``plan.json``.
|
||||
"""
|
||||
self._session_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Validate clusters — drop those without evidence
|
||||
surviving, dropped = _validate_clusters(clusters)
|
||||
|
||||
# Build the user prompt with diagnosis and cluster info
|
||||
cluster_json = json.dumps(
|
||||
[c.model_dump() for c in surviving],
|
||||
indent=2,
|
||||
default=str,
|
||||
)
|
||||
user_prompt = (
|
||||
f"## Diagnosis\n\n{diagnosis_md}\n\n"
|
||||
f"## Surviving Failure Clusters\n\n```json\n{cluster_json}\n```\n\n"
|
||||
"Propose edits for these clusters. Respond with ONLY JSON: "
|
||||
'{"edits": [...]}'
|
||||
)
|
||||
|
||||
# Make the teacher call
|
||||
from openjarvis.core.types import Message, Role
|
||||
|
||||
messages = [
|
||||
Message(role=Role.SYSTEM, content=_PLANNER_SYSTEM_PROMPT),
|
||||
Message(role=Role.USER, content=user_prompt),
|
||||
]
|
||||
result = self._engine.generate(
|
||||
messages=messages,
|
||||
model=self._model,
|
||||
max_tokens=4096,
|
||||
)
|
||||
|
||||
cost_usd = result.get("cost_usd", 0.0)
|
||||
content = result.get("content", "")
|
||||
|
||||
# Persist teacher trace
|
||||
self._persist_trace(content, cost_usd, result)
|
||||
|
||||
# Parse edits from response
|
||||
edits = self._parse_edits(content)
|
||||
|
||||
# Post-process: assign tiers deterministically
|
||||
edits = assign_tiers(edits)
|
||||
|
||||
# Post-process: downgrade large patches to replacements
|
||||
edits = [
|
||||
maybe_downgrade_to_replace(e, prompt_reader=self._prompt_reader)
|
||||
for e in edits
|
||||
]
|
||||
|
||||
# Wire clusters ↔ edits
|
||||
surviving = self._wire_cluster_edit_ids(surviving, edits)
|
||||
|
||||
# Build the plan
|
||||
all_clusters = surviving + dropped
|
||||
plan = LearningPlan(
|
||||
session_id=self._session_id,
|
||||
diagnosis_summary=diagnosis_md,
|
||||
failure_clusters=all_clusters,
|
||||
edits=edits,
|
||||
teacher_model=self._model,
|
||||
estimated_cost_usd=cost_usd,
|
||||
created_at=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
# Persist plan.json
|
||||
plan_path = self._session_dir / "plan.json"
|
||||
plan_path.write_text(plan.model_dump_json(indent=2), encoding="utf-8")
|
||||
|
||||
return plan
|
||||
|
||||
def _parse_edits(self, content: str) -> list[Edit]:
|
||||
"""Parse edits from the teacher's JSON response."""
|
||||
try:
|
||||
data = json.loads(content)
|
||||
except json.JSONDecodeError:
|
||||
# Try to find JSON in the content
|
||||
import re
|
||||
|
||||
match = re.search(r"\{[\s\S]*\}", content)
|
||||
if match:
|
||||
try:
|
||||
data = json.loads(match.group(0))
|
||||
except json.JSONDecodeError:
|
||||
logger.warning("Could not parse edits from teacher response")
|
||||
return []
|
||||
else:
|
||||
logger.warning("No JSON found in teacher response")
|
||||
return []
|
||||
|
||||
raw_edits = data.get("edits", [])
|
||||
edits = []
|
||||
for raw in raw_edits:
|
||||
try:
|
||||
edits.append(Edit.model_validate(raw))
|
||||
except Exception as e:
|
||||
logger.warning("Skipping invalid edit: %s — %s", raw.get("id", "?"), e)
|
||||
return edits
|
||||
|
||||
def _wire_cluster_edit_ids(
|
||||
self,
|
||||
clusters: list[FailureCluster],
|
||||
edits: list[Edit],
|
||||
) -> list[FailureCluster]:
|
||||
"""Populate addressed_by_edit_ids on each cluster."""
|
||||
cluster_map: dict[str, list[str]] = {c.id: [] for c in clusters}
|
||||
for edit in edits:
|
||||
if edit.expected_improvement in cluster_map:
|
||||
cluster_map[edit.expected_improvement].append(edit.id)
|
||||
return [
|
||||
c.model_copy(update={"addressed_by_edit_ids": cluster_map.get(c.id, [])})
|
||||
for c in clusters
|
||||
]
|
||||
|
||||
def _persist_trace(self, content: str, cost_usd: float, result: dict) -> None:
|
||||
"""Write the planner call to teacher_traces/plan.jsonl."""
|
||||
traces_dir = self._session_dir / "teacher_traces"
|
||||
traces_dir.mkdir(parents=True, exist_ok=True)
|
||||
record = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"call": "planner",
|
||||
"content_length": len(content),
|
||||
"cost_usd": cost_usd,
|
||||
"tokens": result.get("usage", {}).get("total_tokens", 0),
|
||||
}
|
||||
jsonl_path = traces_dir / "plan.jsonl"
|
||||
with jsonl_path.open("a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(record) + "\n")
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Prompt diff utilities for the plan phase.
|
||||
|
||||
Handles the PATCH_SYSTEM_PROMPT → REPLACE_SYSTEM_PROMPT downgrade logic.
|
||||
When the teacher proposes a PATCH edit, the planner checks if the diff
|
||||
would change more than 50% of lines. If so, it downgrades to a full
|
||||
REPLACE so the user sees the complete new prompt in the review queue.
|
||||
|
||||
See spec §6.3.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Callable, Optional
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Threshold: if more than this fraction of lines change, downgrade to REPLACE
|
||||
_DOWNGRADE_THRESHOLD = 0.5
|
||||
|
||||
|
||||
def changed_line_ratio(original: str, modified: str) -> float:
|
||||
"""Compute the fraction of lines that differ between two strings.
|
||||
|
||||
Uses a simple line-by-line comparison. Returns 0.0 if both are empty,
|
||||
1.0 if one is empty and the other is not.
|
||||
"""
|
||||
orig_lines = original.splitlines()
|
||||
mod_lines = modified.splitlines()
|
||||
|
||||
if not orig_lines and not mod_lines:
|
||||
return 0.0
|
||||
if not orig_lines or not mod_lines:
|
||||
return 1.0
|
||||
|
||||
max_len = max(len(orig_lines), len(mod_lines))
|
||||
changed = 0
|
||||
for i in range(max_len):
|
||||
orig = orig_lines[i] if i < len(orig_lines) else None
|
||||
mod = mod_lines[i] if i < len(mod_lines) else None
|
||||
if orig != mod:
|
||||
changed += 1
|
||||
|
||||
return changed / max_len
|
||||
|
||||
|
||||
def apply_unified_diff(original: str, diff: str) -> Optional[str]:
|
||||
"""Apply a unified diff to the original string.
|
||||
|
||||
Returns the patched string, or None if the diff cannot be applied.
|
||||
This is a simplified implementation that handles basic unified diffs.
|
||||
"""
|
||||
try:
|
||||
lines = original.splitlines(keepends=True)
|
||||
result_lines: list[str] = []
|
||||
diff_lines = diff.splitlines(keepends=True)
|
||||
|
||||
# Skip header lines (--- and +++)
|
||||
i = 0
|
||||
while i < len(diff_lines) and not diff_lines[i].startswith("@@"):
|
||||
i += 1
|
||||
|
||||
# No hunk headers found — not a valid unified diff
|
||||
if i >= len(diff_lines):
|
||||
return None
|
||||
|
||||
# Parse hunks
|
||||
src_idx = 0
|
||||
while i < len(diff_lines):
|
||||
line = diff_lines[i]
|
||||
if line.startswith("@@"):
|
||||
# Parse hunk header: @@ -start,count +start,count @@
|
||||
match = re.match(r"@@ -(\d+)", line)
|
||||
if not match:
|
||||
return None
|
||||
hunk_start = int(match.group(1)) - 1 # 0-indexed
|
||||
# Copy lines before this hunk
|
||||
while src_idx < hunk_start:
|
||||
if src_idx < len(lines):
|
||||
result_lines.append(lines[src_idx])
|
||||
src_idx += 1
|
||||
i += 1
|
||||
continue
|
||||
|
||||
if line.startswith("-"):
|
||||
# Remove line from original
|
||||
src_idx += 1
|
||||
elif line.startswith("+"):
|
||||
# Add line to result
|
||||
content = line[1:]
|
||||
if not content.endswith("\n"):
|
||||
content += "\n"
|
||||
result_lines.append(content)
|
||||
elif line.startswith(" "):
|
||||
# Context line — copy from original
|
||||
if src_idx < len(lines):
|
||||
result_lines.append(lines[src_idx])
|
||||
src_idx += 1
|
||||
i += 1
|
||||
|
||||
# Copy remaining lines
|
||||
while src_idx < len(lines):
|
||||
result_lines.append(lines[src_idx])
|
||||
src_idx += 1
|
||||
|
||||
return "".join(result_lines)
|
||||
except Exception:
|
||||
logger.warning("Failed to apply unified diff")
|
||||
return None
|
||||
|
||||
|
||||
def maybe_downgrade_to_replace(
|
||||
edit: Edit,
|
||||
*,
|
||||
prompt_reader: Callable[[str], str],
|
||||
) -> Edit:
|
||||
"""Downgrade PATCH_SYSTEM_PROMPT to REPLACE if the diff is large.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
edit :
|
||||
The edit to check.
|
||||
prompt_reader :
|
||||
A callable that takes a target string (e.g. "agents.simple.system_prompt")
|
||||
and returns the current prompt content.
|
||||
|
||||
Returns the edit unchanged if it's not a PATCH op, or if the diff is
|
||||
small enough. Returns a new REPLACE edit if the diff changes > 50% of
|
||||
lines or if the diff cannot be applied.
|
||||
"""
|
||||
if edit.op != EditOp.PATCH_SYSTEM_PROMPT:
|
||||
return edit
|
||||
|
||||
diff_str = edit.payload.get("diff", "")
|
||||
original = prompt_reader(edit.target)
|
||||
patched = apply_unified_diff(original, diff_str)
|
||||
|
||||
if patched is None:
|
||||
# Can't apply the diff — downgrade to REPLACE with a warning
|
||||
logger.warning(
|
||||
"Edit %s: diff could not be applied, downgrading to REPLACE",
|
||||
edit.id,
|
||||
)
|
||||
return edit.model_copy(
|
||||
update={
|
||||
"op": EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
"payload": {"new_content": diff_str},
|
||||
}
|
||||
)
|
||||
|
||||
ratio = changed_line_ratio(original, patched)
|
||||
if ratio > _DOWNGRADE_THRESHOLD:
|
||||
logger.info(
|
||||
"Edit %s: diff changes %.0f%% of lines (>%.0f%%), "
|
||||
"downgrading PATCH → REPLACE",
|
||||
edit.id,
|
||||
ratio * 100,
|
||||
_DOWNGRADE_THRESHOLD * 100,
|
||||
)
|
||||
return edit.model_copy(
|
||||
update={
|
||||
"op": EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
"payload": {"new_content": patched},
|
||||
}
|
||||
)
|
||||
|
||||
return edit
|
||||
@@ -0,0 +1,63 @@
|
||||
"""Deterministic risk tier assignment for edits.
|
||||
|
||||
The teacher cannot pick its own tier. After the teacher emits edits, the
|
||||
planner overwrites each edit's ``risk_tier`` from the lookup table below.
|
||||
If the teacher attempted a different tier, it is silently overwritten and
|
||||
the discrepancy is logged but not surfaced as an error.
|
||||
|
||||
See spec §4.1 (tier table) and §6.2.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Sequence
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp, EditRiskTier
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# The canonical (op) → tier mapping. Every EditOp must appear here.
|
||||
TIER_TABLE: dict[EditOp, EditRiskTier] = {
|
||||
# Intelligence — safe, reversible
|
||||
EditOp.SET_MODEL_FOR_QUERY_CLASS: EditRiskTier.AUTO,
|
||||
EditOp.SET_MODEL_PARAM: EditRiskTier.AUTO,
|
||||
# Agent — params are safe, prompts and class need review
|
||||
EditOp.PATCH_SYSTEM_PROMPT: EditRiskTier.REVIEW,
|
||||
EditOp.REPLACE_SYSTEM_PROMPT: EditRiskTier.REVIEW,
|
||||
EditOp.SET_AGENT_CLASS: EditRiskTier.REVIEW,
|
||||
EditOp.SET_AGENT_PARAM: EditRiskTier.AUTO,
|
||||
EditOp.EDIT_FEW_SHOT_EXEMPLARS: EditRiskTier.REVIEW,
|
||||
# Tools — all safe, reversible
|
||||
EditOp.ADD_TOOL_TO_AGENT: EditRiskTier.AUTO,
|
||||
EditOp.REMOVE_TOOL_FROM_AGENT: EditRiskTier.AUTO,
|
||||
EditOp.EDIT_TOOL_DESCRIPTION: EditRiskTier.AUTO,
|
||||
# v2 — always manual
|
||||
EditOp.LORA_FINETUNE: EditRiskTier.MANUAL,
|
||||
}
|
||||
|
||||
|
||||
def assign_tier(op: EditOp) -> EditRiskTier:
|
||||
"""Return the deterministic risk tier for a given edit op."""
|
||||
return TIER_TABLE[op]
|
||||
|
||||
|
||||
def assign_tiers(edits: Sequence[Edit]) -> list[Edit]:
|
||||
"""Overwrite each edit's risk_tier from the canonical lookup table.
|
||||
|
||||
Returns a new list of Edit objects (pydantic copies). If the teacher
|
||||
had a different tier, it is silently overwritten and logged.
|
||||
"""
|
||||
result = []
|
||||
for edit in edits:
|
||||
correct_tier = assign_tier(edit.op)
|
||||
if edit.risk_tier != correct_tier:
|
||||
logger.info(
|
||||
"Edit %s: overwriting teacher tier %s → %s (op=%s)",
|
||||
edit.id,
|
||||
edit.risk_tier.value,
|
||||
correct_tier.value,
|
||||
edit.op.value,
|
||||
)
|
||||
result.append(edit.model_copy(update={"risk_tier": correct_tier}))
|
||||
return result
|
||||
@@ -0,0 +1 @@
|
||||
"""Storage primitives: paths, SQLite session store."""
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Filesystem path resolution for the distillation subsystem.
|
||||
|
||||
The keystone of artifact isolation (spec §11): the resolved distillation root
|
||||
must NEVER be inside the OpenJarvis source tree. ``resolve_distillation_root``
|
||||
walks up from this module's ``__file__`` looking for a ``pyproject.toml`` that
|
||||
identifies the OpenJarvis source root, then refuses to operate if the resolved
|
||||
root is inside it. Defense in depth — if a user accidentally points
|
||||
``OPENJARVIS_HOME`` at the repo, the system fails loudly instead of silently
|
||||
writing artifacts into the working tree.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.security.file_utils import secure_mkdir
|
||||
|
||||
|
||||
class ConfigurationError(RuntimeError):
|
||||
"""Raised when path configuration would violate isolation guarantees."""
|
||||
|
||||
|
||||
def _find_source_root() -> Path | None:
|
||||
"""Walk upward from this module to find the OpenJarvis source root.
|
||||
|
||||
Returns the directory containing the OpenJarvis ``pyproject.toml``, or
|
||||
``None`` if no such file is found (e.g. when running from an installed
|
||||
wheel rather than a source checkout).
|
||||
"""
|
||||
here = Path(__file__).resolve()
|
||||
for candidate in (here, *here.parents):
|
||||
py = candidate / "pyproject.toml"
|
||||
if py.exists():
|
||||
try:
|
||||
content = py.read_text(encoding="utf-8")
|
||||
except OSError:
|
||||
continue
|
||||
if 'name = "openjarvis"' in content.lower():
|
||||
return candidate
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_openjarvis_home() -> Path:
|
||||
"""Resolve the OPENJARVIS_HOME directory (env var or default)."""
|
||||
env = os.environ.get("OPENJARVIS_HOME")
|
||||
if env:
|
||||
return Path(env).expanduser().resolve()
|
||||
return (Path.home() / ".openjarvis").resolve()
|
||||
|
||||
|
||||
def resolve_distillation_root() -> Path:
|
||||
"""Return the absolute path of the distillation root directory.
|
||||
|
||||
The root is ``$OPENJARVIS_HOME/learning`` (or ``~/.openjarvis/learning``
|
||||
by default). Raises ``ConfigurationError`` if the resolved path lies
|
||||
inside the OpenJarvis source tree, to prevent dev artifacts from leaking
|
||||
into the repo.
|
||||
"""
|
||||
home = _resolve_openjarvis_home()
|
||||
source_root = _find_source_root()
|
||||
if source_root is not None:
|
||||
try:
|
||||
home.relative_to(source_root)
|
||||
except ValueError:
|
||||
pass # Good — not inside the source tree.
|
||||
else:
|
||||
raise ConfigurationError(
|
||||
f"OPENJARVIS_HOME ({home}) is inside the source tree "
|
||||
f"({source_root}). Distillation refuses to write runtime "
|
||||
"artifacts inside the OpenJarvis repo. Set OPENJARVIS_HOME "
|
||||
"to a directory outside the repo (default: ~/.openjarvis)."
|
||||
)
|
||||
return home / "learning"
|
||||
|
||||
|
||||
def ensure_distillation_dirs() -> Path:
|
||||
"""Create the distillation directory layout if missing.
|
||||
|
||||
Returns the distillation root. Creates ``sessions/``, ``benchmarks/``,
|
||||
``benchmarks/reference_outputs/``, and ``pending_review/`` underneath it,
|
||||
all with restrictive ``0o700`` permissions via ``secure_mkdir``.
|
||||
"""
|
||||
root = resolve_distillation_root()
|
||||
secure_mkdir(root)
|
||||
secure_mkdir(root / "sessions")
|
||||
secure_mkdir(root / "benchmarks")
|
||||
secure_mkdir(root / "benchmarks" / "reference_outputs")
|
||||
secure_mkdir(root / "pending_review")
|
||||
return root
|
||||
@@ -0,0 +1,304 @@
|
||||
"""SQLite-backed storage for distillation LearningSession records.
|
||||
|
||||
Mirrors the style of ``openjarvis.learning.optimize.store.OptimizationStore``:
|
||||
|
||||
- stdlib ``sqlite3`` in WAL mode
|
||||
- inline DDL as module-level constants
|
||||
- persistent connection stored as ``self._conn``
|
||||
- ``_migrate()`` runs additive ALTER TABLEs that swallow ``OperationalError``
|
||||
|
||||
The store does NOT share its database file with ``OptimizationStore`` (see
|
||||
spec §8.1 / brainstorming Q8). The two SQLite files live side-by-side in
|
||||
``~/.openjarvis/learning/`` but are independent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import sqlite3
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Optional, Union
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
BenchmarkSnapshot,
|
||||
EditOutcome,
|
||||
LearningSession,
|
||||
SessionStatus,
|
||||
TriggerKind,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
_CREATE_SESSIONS = """\
|
||||
CREATE TABLE IF NOT EXISTS learning_sessions (
|
||||
id TEXT PRIMARY KEY,
|
||||
parent_session_id TEXT,
|
||||
trigger TEXT NOT NULL,
|
||||
trigger_metadata TEXT NOT NULL DEFAULT '{}',
|
||||
status TEXT NOT NULL,
|
||||
autonomy_mode TEXT NOT NULL,
|
||||
started_at TEXT NOT NULL,
|
||||
ended_at TEXT,
|
||||
diagnosis_path TEXT NOT NULL,
|
||||
plan_path TEXT NOT NULL,
|
||||
benchmark_before TEXT NOT NULL,
|
||||
benchmark_after TEXT,
|
||||
git_checkpoint_pre TEXT NOT NULL,
|
||||
git_checkpoint_post TEXT,
|
||||
teacher_cost_usd REAL NOT NULL DEFAULT 0.0,
|
||||
error TEXT,
|
||||
FOREIGN KEY (parent_session_id) REFERENCES learning_sessions(id)
|
||||
);
|
||||
"""
|
||||
|
||||
_CREATE_OUTCOMES = """\
|
||||
CREATE TABLE IF NOT EXISTS edit_outcomes (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
session_id TEXT NOT NULL,
|
||||
edit_id TEXT NOT NULL,
|
||||
pillar TEXT NOT NULL,
|
||||
op TEXT NOT NULL,
|
||||
target TEXT NOT NULL,
|
||||
risk_tier TEXT NOT NULL,
|
||||
status TEXT NOT NULL,
|
||||
benchmark_delta REAL,
|
||||
cluster_deltas TEXT NOT NULL DEFAULT '{}',
|
||||
rationale TEXT NOT NULL DEFAULT '',
|
||||
error TEXT,
|
||||
applied_at TEXT,
|
||||
FOREIGN KEY (session_id) REFERENCES learning_sessions(id)
|
||||
);
|
||||
"""
|
||||
|
||||
_CREATE_INDEXES = [
|
||||
"CREATE INDEX IF NOT EXISTS idx_sessions_started_at "
|
||||
"ON learning_sessions(started_at)",
|
||||
"CREATE INDEX IF NOT EXISTS idx_sessions_status ON learning_sessions(status)",
|
||||
"CREATE INDEX IF NOT EXISTS idx_outcomes_session ON edit_outcomes(session_id)",
|
||||
"CREATE INDEX IF NOT EXISTS idx_outcomes_op ON edit_outcomes(op)",
|
||||
]
|
||||
|
||||
_INSERT_SESSION = """\
|
||||
INSERT OR REPLACE INTO learning_sessions (
|
||||
id, parent_session_id, trigger, trigger_metadata, status, autonomy_mode,
|
||||
started_at, ended_at, diagnosis_path, plan_path,
|
||||
benchmark_before, benchmark_after,
|
||||
git_checkpoint_pre, git_checkpoint_post, teacher_cost_usd, error
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"""
|
||||
|
||||
_INSERT_OUTCOME = """\
|
||||
INSERT INTO edit_outcomes (
|
||||
session_id, edit_id, pillar, op, target, risk_tier, status,
|
||||
benchmark_delta, cluster_deltas, rationale, error, applied_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"""
|
||||
|
||||
# Future ALTER TABLE statements go here, swallowing OperationalError.
|
||||
_MIGRATE: list[str] = []
|
||||
|
||||
|
||||
def _dt_to_iso(dt: datetime | None) -> str | None:
|
||||
return dt.isoformat() if dt is not None else None
|
||||
|
||||
|
||||
def _iso_to_dt(s: str | None) -> datetime | None:
|
||||
return datetime.fromisoformat(s) if s else None
|
||||
|
||||
|
||||
class SessionStore:
|
||||
"""SQLite-backed storage for LearningSession and EditOutcome records.
|
||||
|
||||
The full LearningSession is also serialized to disk as
|
||||
``<session_dir>/session.json`` — that file is the authoritative source
|
||||
if SQLite is ever lost. This store is the index used for fast queries.
|
||||
"""
|
||||
|
||||
def __init__(self, db_path: Union[str, Path]) -> None:
|
||||
self._db_path = str(db_path)
|
||||
self._conn = sqlite3.connect(self._db_path, check_same_thread=False)
|
||||
self._conn.execute("PRAGMA journal_mode=WAL")
|
||||
self._conn.execute("PRAGMA foreign_keys=ON")
|
||||
self._conn.execute(_CREATE_SESSIONS)
|
||||
self._conn.execute(_CREATE_OUTCOMES)
|
||||
for index in _CREATE_INDEXES:
|
||||
self._conn.execute(index)
|
||||
self._conn.commit()
|
||||
self._migrate()
|
||||
|
||||
def _migrate(self) -> None:
|
||||
"""Apply additive schema migrations, swallowing already-applied ones."""
|
||||
for stmt in _MIGRATE:
|
||||
try:
|
||||
self._conn.execute(stmt)
|
||||
except sqlite3.OperationalError:
|
||||
pass
|
||||
self._conn.commit()
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the underlying SQLite connection."""
|
||||
self._conn.close()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Sessions
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def save_session(self, session: LearningSession) -> None:
|
||||
"""Insert or update a LearningSession (idempotent on session.id)."""
|
||||
self._conn.execute(
|
||||
_INSERT_SESSION,
|
||||
(
|
||||
session.id,
|
||||
session.parent_session_id,
|
||||
session.trigger.value,
|
||||
json.dumps(session.trigger_metadata),
|
||||
session.status.value,
|
||||
session.autonomy_mode.value,
|
||||
session.started_at.isoformat(),
|
||||
_dt_to_iso(session.ended_at),
|
||||
str(session.diagnosis_path),
|
||||
str(session.plan_path),
|
||||
session.benchmark_before.model_dump_json(),
|
||||
(
|
||||
session.benchmark_after.model_dump_json()
|
||||
if session.benchmark_after is not None
|
||||
else None
|
||||
),
|
||||
session.git_checkpoint_pre,
|
||||
session.git_checkpoint_post,
|
||||
session.teacher_cost_usd,
|
||||
session.error,
|
||||
),
|
||||
)
|
||||
self._conn.commit()
|
||||
|
||||
def get_session(self, session_id: str) -> Optional[LearningSession]:
|
||||
"""Return the LearningSession with the given id, or None if missing."""
|
||||
row = self._conn.execute(
|
||||
"SELECT id, parent_session_id, trigger, trigger_metadata, status, "
|
||||
"autonomy_mode, started_at, ended_at, diagnosis_path, plan_path, "
|
||||
"benchmark_before, benchmark_after, git_checkpoint_pre, "
|
||||
"git_checkpoint_post, teacher_cost_usd, error "
|
||||
"FROM learning_sessions WHERE id = ?",
|
||||
(session_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
return self._row_to_session(row, with_outcomes=True)
|
||||
|
||||
def list_sessions(
|
||||
self,
|
||||
status: SessionStatus | None = None,
|
||||
limit: int | None = None,
|
||||
) -> list[LearningSession]:
|
||||
"""List sessions ordered by ``started_at DESC``."""
|
||||
sql = (
|
||||
"SELECT id, parent_session_id, trigger, trigger_metadata, status, "
|
||||
"autonomy_mode, started_at, ended_at, diagnosis_path, plan_path, "
|
||||
"benchmark_before, benchmark_after, git_checkpoint_pre, "
|
||||
"git_checkpoint_post, teacher_cost_usd, error "
|
||||
"FROM learning_sessions"
|
||||
)
|
||||
params: list[object] = []
|
||||
if status is not None:
|
||||
sql += " WHERE status = ?"
|
||||
params.append(status.value)
|
||||
sql += " ORDER BY started_at DESC"
|
||||
if limit is not None:
|
||||
sql += " LIMIT ?"
|
||||
params.append(limit)
|
||||
rows = self._conn.execute(sql, params).fetchall()
|
||||
return [self._row_to_session(r, with_outcomes=True) for r in rows]
|
||||
|
||||
def _row_to_session(self, row: tuple, with_outcomes: bool) -> LearningSession:
|
||||
session = LearningSession(
|
||||
id=row[0],
|
||||
parent_session_id=row[1],
|
||||
trigger=TriggerKind(row[2]),
|
||||
trigger_metadata=json.loads(row[3]),
|
||||
status=SessionStatus(row[4]),
|
||||
autonomy_mode=AutonomyMode(row[5]),
|
||||
started_at=datetime.fromisoformat(row[6]),
|
||||
ended_at=_iso_to_dt(row[7]),
|
||||
diagnosis_path=Path(row[8]),
|
||||
plan_path=Path(row[9]),
|
||||
benchmark_before=BenchmarkSnapshot.model_validate_json(row[10]),
|
||||
benchmark_after=(
|
||||
BenchmarkSnapshot.model_validate_json(row[11])
|
||||
if row[11] is not None
|
||||
else None
|
||||
),
|
||||
edit_outcomes=[],
|
||||
git_checkpoint_pre=row[12],
|
||||
git_checkpoint_post=row[13],
|
||||
teacher_cost_usd=row[14],
|
||||
error=row[15],
|
||||
)
|
||||
if with_outcomes:
|
||||
outcomes = self.list_outcomes(session.id)
|
||||
session = session.model_copy(update={"edit_outcomes": outcomes})
|
||||
return session
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Edit outcomes
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def save_outcome(
|
||||
self,
|
||||
session_id: str,
|
||||
outcome: EditOutcome,
|
||||
*,
|
||||
pillar: str,
|
||||
op: str,
|
||||
target: str,
|
||||
risk_tier: str,
|
||||
rationale: str = "",
|
||||
) -> None:
|
||||
"""Insert an EditOutcome row.
|
||||
|
||||
``pillar``, ``op``, ``target``, ``risk_tier``, and ``rationale`` come
|
||||
from the parent ``Edit`` (which is not stored on the EditOutcome
|
||||
model). They are kept as columns to make ``WHERE op = ?`` queries
|
||||
possible without joining against the on-disk plan.json.
|
||||
"""
|
||||
self._conn.execute(
|
||||
_INSERT_OUTCOME,
|
||||
(
|
||||
session_id,
|
||||
outcome.edit_id,
|
||||
pillar,
|
||||
op,
|
||||
target,
|
||||
risk_tier,
|
||||
outcome.status,
|
||||
outcome.benchmark_delta,
|
||||
json.dumps(outcome.cluster_deltas),
|
||||
rationale,
|
||||
outcome.error,
|
||||
_dt_to_iso(outcome.applied_at),
|
||||
),
|
||||
)
|
||||
self._conn.commit()
|
||||
|
||||
def list_outcomes(self, session_id: str) -> list[EditOutcome]:
|
||||
"""Return all EditOutcomes for a session, ordered by insertion id."""
|
||||
rows = self._conn.execute(
|
||||
"SELECT edit_id, status, benchmark_delta, cluster_deltas, error, "
|
||||
"applied_at FROM edit_outcomes WHERE session_id = ? ORDER BY id",
|
||||
(session_id,),
|
||||
).fetchall()
|
||||
return [
|
||||
EditOutcome(
|
||||
edit_id=row[0],
|
||||
status=row[1],
|
||||
benchmark_delta=row[2],
|
||||
cluster_deltas=json.loads(row[3]),
|
||||
error=row[4],
|
||||
applied_at=_iso_to_dt(row[5]),
|
||||
)
|
||||
for row in rows
|
||||
]
|
||||
@@ -0,0 +1,135 @@
|
||||
"""Real student runner for distillation experiments.
|
||||
|
||||
Replaces the ``MagicMock()`` in the experiment runner script with a
|
||||
callable that actually invokes the student model via vLLM (or any
|
||||
OpenAI-compatible engine) and returns structured results.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class StudentResult:
|
||||
"""Result from running the student model on a task."""
|
||||
|
||||
content: str
|
||||
score: float = 0.0
|
||||
trace_id: str = ""
|
||||
latency_seconds: float = 0.0
|
||||
tokens_used: int = 0
|
||||
|
||||
|
||||
class VLLMStudentRunner:
|
||||
"""Invoke the student model via a vLLM OpenAI-compatible endpoint.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
host :
|
||||
vLLM server URL (e.g. ``http://localhost:8001``).
|
||||
model :
|
||||
Model name as registered in vLLM (e.g. ``Qwen/Qwen3.5-9B``).
|
||||
temperature :
|
||||
Sampling temperature.
|
||||
max_tokens :
|
||||
Max tokens for the student response.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str = "http://localhost:8001",
|
||||
model: str = "Qwen/Qwen3.5-9B",
|
||||
temperature: float = 0.6,
|
||||
max_tokens: int = 4096,
|
||||
) -> None:
|
||||
import httpx
|
||||
|
||||
self._host = host.rstrip("/")
|
||||
self._model = model
|
||||
self._temperature = temperature
|
||||
self._max_tokens = max_tokens
|
||||
self._client = httpx.Client(base_url=self._host, timeout=300.0)
|
||||
|
||||
def __call__(
|
||||
self, query: str, session_id: str = "", **kwargs: Any
|
||||
) -> StudentResult:
|
||||
"""Run the student on *query* and return a StudentResult."""
|
||||
t0 = time.time()
|
||||
try:
|
||||
resp = self._client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": self._model,
|
||||
"messages": [{"role": "user", "content": query}],
|
||||
"temperature": self._temperature,
|
||||
"max_tokens": self._max_tokens,
|
||||
"stream": False,
|
||||
},
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except Exception as exc:
|
||||
latency = time.time() - t0
|
||||
logger.warning("Student runner failed: %s", exc)
|
||||
return StudentResult(
|
||||
content=f"Error: {exc}",
|
||||
latency_seconds=latency,
|
||||
)
|
||||
|
||||
latency = time.time() - t0
|
||||
choices = data.get("choices", [])
|
||||
content = ""
|
||||
if choices:
|
||||
msg = choices[0].get("message", {})
|
||||
content = msg.get("content", "")
|
||||
|
||||
usage = data.get("usage", {})
|
||||
tokens = usage.get("total_tokens", 0)
|
||||
|
||||
return StudentResult(
|
||||
content=content,
|
||||
latency_seconds=latency,
|
||||
tokens_used=tokens,
|
||||
trace_id=f"distill_{session_id}_{hash(query) % 10000}",
|
||||
)
|
||||
|
||||
|
||||
def build_benchmark_samples_from_traces(
|
||||
trace_store: Any,
|
||||
*,
|
||||
limit: int = 50,
|
||||
min_feedback: float | None = None,
|
||||
) -> list:
|
||||
"""Build PersonalBenchmarkSample objects from the trace store.
|
||||
|
||||
Pulls recent traces (optionally filtered by feedback score) and
|
||||
converts them into benchmark samples the teacher can reference.
|
||||
"""
|
||||
from openjarvis.learning.optimize.personal.synthesizer import (
|
||||
PersonalBenchmarkSample,
|
||||
)
|
||||
|
||||
traces = trace_store.list_traces(limit=limit)
|
||||
samples = []
|
||||
for t in traces:
|
||||
fb = getattr(t, "feedback", None)
|
||||
if min_feedback is not None and (fb is None or fb < min_feedback):
|
||||
continue
|
||||
samples.append(
|
||||
PersonalBenchmarkSample(
|
||||
trace_id=t.trace_id,
|
||||
query=t.query,
|
||||
reference_answer=t.result[:2000] if t.result else "",
|
||||
agent=t.agent,
|
||||
category="benchmark",
|
||||
feedback_score=fb if fb is not None else 0.0,
|
||||
)
|
||||
)
|
||||
logger.info("Built %d benchmark samples from traces", len(samples))
|
||||
return samples
|
||||
@@ -0,0 +1,66 @@
|
||||
"""Trigger types for the distillation subsystem.
|
||||
|
||||
A trigger is what kicks off a learning session. Four trigger types exist,
|
||||
all funneling into ``DistillationOrchestrator.run(trigger)``. The trigger
|
||||
object is stored on the ``LearningSession`` for queryability.
|
||||
|
||||
See spec §3.3.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from openjarvis.learning.distillation.models import TriggerKind
|
||||
|
||||
|
||||
@dataclass
|
||||
class OnDemandTrigger:
|
||||
"""User ran ``jarvis learning run`` from the CLI."""
|
||||
|
||||
kind: TriggerKind = TriggerKind.ON_DEMAND
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserFlagTrigger:
|
||||
"""User flagged a specific trace for improvement."""
|
||||
|
||||
trace_id: str = ""
|
||||
kind: TriggerKind = TriggerKind.USER_FLAG
|
||||
|
||||
@property
|
||||
def metadata(self) -> dict[str, Any]:
|
||||
return {"trace_id": self.trace_id}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScheduledTrigger:
|
||||
"""Cron-based scheduled trigger."""
|
||||
|
||||
cron: str = "0 3 * * *"
|
||||
new_trace_count: int = 0
|
||||
kind: TriggerKind = TriggerKind.SCHEDULED
|
||||
|
||||
@property
|
||||
def metadata(self) -> dict[str, Any]:
|
||||
return {"cron": self.cron, "new_trace_count": self.new_trace_count}
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClusterTrigger:
|
||||
"""Fired when a failure cluster exceeds a threshold."""
|
||||
|
||||
cluster_description: str = ""
|
||||
trace_ids: list[str] = field(default_factory=list)
|
||||
failure_rate: float = 0.0
|
||||
kind: TriggerKind = TriggerKind.CLUSTER
|
||||
|
||||
@property
|
||||
def metadata(self) -> dict[str, Any]:
|
||||
return {
|
||||
"cluster_description": self.cluster_description,
|
||||
"trace_ids": self.trace_ids,
|
||||
"failure_rate": self.failure_rate,
|
||||
}
|
||||
@@ -64,12 +64,17 @@ class BaseTool(ABC):
|
||||
|
||||
def to_openai_function(self) -> Dict[str, Any]:
|
||||
"""Convert to OpenAI function-calling format."""
|
||||
from openjarvis.tools.description_loader import (
|
||||
get_tool_description_override,
|
||||
)
|
||||
|
||||
s = self.spec
|
||||
desc = get_tool_description_override(s.name) or s.description
|
||||
return {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": s.name,
|
||||
"description": s.description,
|
||||
"description": desc,
|
||||
"parameters": s.parameters,
|
||||
},
|
||||
}
|
||||
@@ -356,10 +361,15 @@ def build_tool_descriptions(
|
||||
if not tools:
|
||||
return "No tools available."
|
||||
|
||||
from openjarvis.tools.description_loader import (
|
||||
get_tool_description_override,
|
||||
)
|
||||
|
||||
sections: list[str] = []
|
||||
for t in tools:
|
||||
s = t.spec
|
||||
lines = [f"### {s.name}", s.description]
|
||||
desc = get_tool_description_override(s.name) or s.description
|
||||
lines = [f"### {s.name}", desc]
|
||||
|
||||
if include_category and s.category:
|
||||
lines.append(f"Category: {s.category}")
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
"""Load tool description overrides from $OPENJARVIS_HOME/tools/descriptions.toml.
|
||||
|
||||
Distillation (M1) proposes tool description edits that get written to disk by
|
||||
``EditToolDescriptionApplier``. This module loads those overrides so agents
|
||||
see the improved descriptions at runtime.
|
||||
|
||||
The TOML file format (written by the applier) is::
|
||||
|
||||
[web_search]
|
||||
description = "Search the web for recent information only"
|
||||
|
||||
[llm]
|
||||
description = "Call a sub-LM. Has no internet access."
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_cache: Optional[Dict[str, str]] = None
|
||||
|
||||
|
||||
def _load_overrides() -> Dict[str, str]:
|
||||
"""Parse descriptions.toml and return {tool_name: description}."""
|
||||
home = Path(os.environ.get("OPENJARVIS_HOME", "~/.openjarvis")).expanduser()
|
||||
desc_path = home / "tools" / "descriptions.toml"
|
||||
if not desc_path.exists():
|
||||
return {}
|
||||
try:
|
||||
content = desc_path.read_text(encoding="utf-8")
|
||||
except Exception:
|
||||
logger.warning(
|
||||
"Failed to read tool description overrides at %s",
|
||||
desc_path,
|
||||
exc_info=True,
|
||||
)
|
||||
return {}
|
||||
|
||||
overrides: Dict[str, str] = {}
|
||||
current_tool: Optional[str] = None
|
||||
for line in content.splitlines():
|
||||
stripped = line.strip()
|
||||
if stripped.startswith("[") and stripped.endswith("]"):
|
||||
current_tool = stripped[1:-1]
|
||||
elif current_tool and stripped.startswith("description"):
|
||||
# Parse: description = "..."
|
||||
_, _, value = stripped.partition("=")
|
||||
value = value.strip().strip('"').strip("'")
|
||||
if value:
|
||||
overrides[current_tool] = value
|
||||
if overrides:
|
||||
logger.info(
|
||||
"Loaded %d tool description overrides from %s",
|
||||
len(overrides),
|
||||
desc_path,
|
||||
)
|
||||
return overrides
|
||||
|
||||
|
||||
def get_tool_description_override(tool_name: str) -> Optional[str]:
|
||||
"""Return the override description for *tool_name*, or ``None``.
|
||||
|
||||
Results are cached for the lifetime of the process.
|
||||
"""
|
||||
global _cache # noqa: PLW0603
|
||||
if _cache is None:
|
||||
_cache = _load_overrides()
|
||||
return _cache.get(tool_name)
|
||||
|
||||
|
||||
def clear_cache() -> None:
|
||||
"""Clear the cached overrides (useful for testing)."""
|
||||
global _cache # noqa: PLW0603
|
||||
_cache = None
|
||||
@@ -355,7 +355,7 @@ class TestConfigBenchmarks:
|
||||
def test_benchmarks_count(self) -> None:
|
||||
from openjarvis.evals.core.config import KNOWN_BENCHMARKS
|
||||
|
||||
assert len(KNOWN_BENCHMARKS) == 30
|
||||
assert len(KNOWN_BENCHMARKS) == 31
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Tests for the distillation subsystem."""
|
||||
@@ -0,0 +1,105 @@
|
||||
"""Tests for openjarvis.learning.distillation.execute.base module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class TestApplyContext:
|
||||
"""Tests for ApplyContext dataclass."""
|
||||
|
||||
def test_constructs(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
assert ctx.openjarvis_home == tmp_path
|
||||
assert ctx.session_id == "s1"
|
||||
|
||||
def test_config_path(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
assert ctx.config_path == tmp_path / "config.toml"
|
||||
|
||||
def test_agents_dir(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
assert ctx.agents_dir == tmp_path / "agents"
|
||||
|
||||
def test_tools_dir(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
assert ctx.tools_dir == tmp_path / "tools"
|
||||
|
||||
|
||||
class TestValidationResult:
|
||||
"""Tests for ValidationResult."""
|
||||
|
||||
def test_ok_result(self) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ValidationResult
|
||||
|
||||
r = ValidationResult(ok=True)
|
||||
assert r.ok is True
|
||||
assert r.reason == ""
|
||||
|
||||
def test_error_result(self) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import ValidationResult
|
||||
|
||||
r = ValidationResult(ok=False, reason="target not found")
|
||||
assert r.ok is False
|
||||
assert r.reason == "target not found"
|
||||
|
||||
|
||||
class TestEditApplierRegistry:
|
||||
"""Tests for EditApplierRegistry."""
|
||||
|
||||
def test_register_and_get(self) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
ApplyContext,
|
||||
ApplyResult,
|
||||
EditApplier,
|
||||
EditApplierRegistry,
|
||||
ValidationResult,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import Edit, EditOp
|
||||
|
||||
class FakeApplier(EditApplier):
|
||||
op = EditOp.SET_MODEL_PARAM
|
||||
|
||||
def validate(self, edit: Edit, ctx: ApplyContext) -> ValidationResult:
|
||||
return ValidationResult(ok=True)
|
||||
|
||||
def apply(self, edit: Edit, ctx: ApplyContext) -> ApplyResult:
|
||||
return ApplyResult()
|
||||
|
||||
def rollback(self, edit: Edit, ctx: ApplyContext) -> None:
|
||||
pass
|
||||
|
||||
registry = EditApplierRegistry()
|
||||
registry.register(FakeApplier())
|
||||
assert registry.is_supported(EditOp.SET_MODEL_PARAM)
|
||||
applier = registry.get(EditOp.SET_MODEL_PARAM)
|
||||
assert isinstance(applier, FakeApplier)
|
||||
|
||||
def test_is_supported_returns_false_for_unregistered(self) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
EditApplierRegistry,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
registry = EditApplierRegistry()
|
||||
assert registry.is_supported(EditOp.LORA_FINETUNE) is False
|
||||
|
||||
def test_get_raises_for_unregistered(self) -> None:
|
||||
from openjarvis.learning.distillation.execute.base import (
|
||||
EditApplierRegistry,
|
||||
)
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
registry = EditApplierRegistry()
|
||||
with pytest.raises(KeyError):
|
||||
registry.get(EditOp.LORA_FINETUNE)
|
||||
@@ -0,0 +1,53 @@
|
||||
"""Tests for LoRA stub applier."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_lora_edit() -> Edit:
|
||||
return Edit(
|
||||
id="edit-lora",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.LORA_FINETUNE,
|
||||
target="models.qwen2.5-coder:7b",
|
||||
payload={"target_model": "qwen2.5-coder:7b", "data_source": "trace_filter:*"},
|
||||
rationale="Fine-tune for math",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.MANUAL,
|
||||
)
|
||||
|
||||
|
||||
class TestLoraStubApplier:
|
||||
"""Tests for LoraStubApplier."""
|
||||
|
||||
def test_validate_returns_not_ok(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.lora_stub import (
|
||||
LoraStubApplier,
|
||||
)
|
||||
|
||||
applier = LoraStubApplier()
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
result = applier.validate(_make_lora_edit(), ctx)
|
||||
assert not result.ok
|
||||
assert "v2" in result.reason.lower() or "deferred" in result.reason.lower()
|
||||
|
||||
def test_apply_raises_not_implemented(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.lora_stub import (
|
||||
LoraStubApplier,
|
||||
)
|
||||
|
||||
applier = LoraStubApplier()
|
||||
ctx = ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
with pytest.raises(NotImplementedError, match="deferred to v2"):
|
||||
applier.apply(_make_lora_edit(), ctx)
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Tests for agent-pillar appliers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_ctx(tmp_path: Path) -> ApplyContext:
|
||||
agents_dir = tmp_path / "agents" / "simple"
|
||||
agents_dir.mkdir(parents=True)
|
||||
(agents_dir / "system_prompt.md").write_text(
|
||||
"You are a helpful assistant.\nBe concise.\n"
|
||||
)
|
||||
(tmp_path / "config.toml").write_text(
|
||||
"[agent]\n"
|
||||
'default = "simple"\n'
|
||||
"\n"
|
||||
"[agent.simple]\n"
|
||||
'class = "simple"\n'
|
||||
"max_turns = 5\n"
|
||||
)
|
||||
return ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
|
||||
|
||||
class TestReplaceSystemPromptApplier:
|
||||
"""Tests for ReplaceSystemPromptApplier."""
|
||||
|
||||
def test_validate_ok(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
ReplaceSystemPromptApplier,
|
||||
)
|
||||
|
||||
applier = ReplaceSystemPromptApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"new_content": "New prompt content.\n"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
assert applier.validate(edit, ctx).ok
|
||||
|
||||
def test_apply_overwrites_prompt(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
ReplaceSystemPromptApplier,
|
||||
)
|
||||
|
||||
applier = ReplaceSystemPromptApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"new_content": "You are a math expert.\n"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = (tmp_path / "agents" / "simple" / "system_prompt.md").read_text()
|
||||
assert content == "You are a math expert.\n"
|
||||
|
||||
|
||||
class TestPatchSystemPromptApplier:
|
||||
"""Tests for PatchSystemPromptApplier."""
|
||||
|
||||
def test_apply_applies_diff(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
PatchSystemPromptApplier,
|
||||
)
|
||||
|
||||
applier = PatchSystemPromptApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
diff = (
|
||||
"--- a/prompt.md\n"
|
||||
"+++ b/prompt.md\n"
|
||||
"@@ -1,2 +1,2 @@\n"
|
||||
" You are a helpful assistant.\n"
|
||||
"-Be concise.\n"
|
||||
"+Be concise and use math tools.\n"
|
||||
)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": diff},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = (tmp_path / "agents" / "simple" / "system_prompt.md").read_text()
|
||||
assert "math tools" in content
|
||||
|
||||
def test_validate_fails_for_bad_diff(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
PatchSystemPromptApplier,
|
||||
)
|
||||
|
||||
applier = PatchSystemPromptApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-002",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": "not a valid diff"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
result = applier.validate(edit, ctx)
|
||||
assert not result.ok
|
||||
|
||||
|
||||
class TestSetAgentClassApplier:
|
||||
"""Tests for SetAgentClassApplier."""
|
||||
|
||||
def test_apply_updates_config(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
SetAgentClassApplier,
|
||||
)
|
||||
|
||||
applier = SetAgentClassApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.SET_AGENT_CLASS,
|
||||
target="agent.simple.class",
|
||||
payload={"agent": "simple", "new_class": "react"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "react" in content
|
||||
|
||||
|
||||
class TestSetAgentParamApplier:
|
||||
"""Tests for SetAgentParamApplier."""
|
||||
|
||||
def test_apply_updates_param(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
SetAgentParamApplier,
|
||||
)
|
||||
|
||||
applier = SetAgentParamApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.SET_AGENT_PARAM,
|
||||
target="agent.simple.max_turns",
|
||||
payload={"agent": "simple", "param": "max_turns", "value": 10},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "10" in content
|
||||
|
||||
|
||||
class TestEditFewShotExemplarsApplier:
|
||||
"""Tests for EditFewShotExemplarsApplier."""
|
||||
|
||||
def test_apply_writes_exemplars(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.agent import (
|
||||
EditFewShotExemplarsApplier,
|
||||
)
|
||||
|
||||
applier = EditFewShotExemplarsApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
exemplars = [
|
||||
{"input": "What is 2+2?", "output": "4"},
|
||||
{"input": "Capital of France?", "output": "Paris"},
|
||||
]
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.EDIT_FEW_SHOT_EXEMPLARS,
|
||||
target="agents.simple.few_shot",
|
||||
payload={"agent": "simple", "exemplars": exemplars},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
fs_path = tmp_path / "agents" / "simple" / "few_shot.json"
|
||||
assert fs_path.exists()
|
||||
data = json.loads(fs_path.read_text())
|
||||
assert len(data) == 2
|
||||
assert data[0]["input"] == "What is 2+2?"
|
||||
@@ -0,0 +1,122 @@
|
||||
"""Tests for intelligence-pillar appliers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_ctx(tmp_path: Path) -> ApplyContext:
|
||||
(tmp_path / "config.toml").write_text(
|
||||
"[learning.routing]\n"
|
||||
'policy = "learned"\n'
|
||||
"\n"
|
||||
"[learning.routing.policy_map]\n"
|
||||
'math = "qwen2.5-coder:3b"\n'
|
||||
'code = "qwen2.5-coder:7b"\n'
|
||||
)
|
||||
return ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
|
||||
|
||||
def _make_routing_edit(
|
||||
query_class: str = "math",
|
||||
model: str = "qwen2.5-coder:14b",
|
||||
) -> Edit:
|
||||
return Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
target="learning.routing.policy_map.math",
|
||||
payload={"query_class": query_class, "model": model},
|
||||
rationale="Route math to bigger model",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
|
||||
|
||||
def _make_param_edit(
|
||||
model: str = "qwen2.5-coder:7b",
|
||||
param: str = "temperature",
|
||||
value: float = 0.3,
|
||||
) -> Edit:
|
||||
return Edit(
|
||||
id="edit-002",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_PARAM,
|
||||
target=f"models.{model}.{param}",
|
||||
payload={"model": model, "param": param, "value": value},
|
||||
rationale="Lower temperature for code",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
|
||||
|
||||
class TestSetModelForQueryClassApplier:
|
||||
"""Tests for SetModelForQueryClassApplier."""
|
||||
|
||||
def test_validate_ok(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelForQueryClassApplier,
|
||||
)
|
||||
|
||||
applier = SetModelForQueryClassApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
result = applier.validate(_make_routing_edit(), ctx)
|
||||
assert result.ok
|
||||
|
||||
def test_apply_updates_config(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelForQueryClassApplier,
|
||||
)
|
||||
|
||||
applier = SetModelForQueryClassApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
applier.apply(_make_routing_edit(), ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "qwen2.5-coder:14b" in content
|
||||
|
||||
def test_apply_adds_new_query_class(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelForQueryClassApplier,
|
||||
)
|
||||
|
||||
applier = SetModelForQueryClassApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = _make_routing_edit(query_class="science", model="qwen2.5-coder:14b")
|
||||
applier.apply(edit, ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "science" in content
|
||||
assert "qwen2.5-coder:14b" in content
|
||||
|
||||
|
||||
class TestSetModelParamApplier:
|
||||
"""Tests for SetModelParamApplier."""
|
||||
|
||||
def test_validate_ok(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelParamApplier,
|
||||
)
|
||||
|
||||
applier = SetModelParamApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
result = applier.validate(_make_param_edit(), ctx)
|
||||
assert result.ok
|
||||
|
||||
def test_apply_writes_param(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.intelligence import (
|
||||
SetModelParamApplier,
|
||||
)
|
||||
|
||||
applier = SetModelParamApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
applier.apply(_make_param_edit(), ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "temperature" in content
|
||||
assert "0.3" in content
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Tests for tools-pillar appliers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_ctx(tmp_path: Path) -> ApplyContext:
|
||||
tools_dir = tmp_path / "tools"
|
||||
tools_dir.mkdir(parents=True)
|
||||
(tools_dir / "descriptions.toml").write_text(
|
||||
'[web_search]\ndescription = "Search the web"\n'
|
||||
)
|
||||
(tmp_path / "config.toml").write_text('[agent.simple]\ntools = ["web_search"]\n')
|
||||
return ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
|
||||
|
||||
class TestAddToolToAgentApplier:
|
||||
"""Tests for AddToolToAgentApplier."""
|
||||
|
||||
def test_apply_adds_tool(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
AddToolToAgentApplier,
|
||||
)
|
||||
|
||||
applier = AddToolToAgentApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.TOOLS,
|
||||
op=EditOp.ADD_TOOL_TO_AGENT,
|
||||
target="agent.simple.tools",
|
||||
payload={"agent": "simple", "tool_name": "calculator"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "calculator" in content
|
||||
|
||||
def test_validate_ok(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
AddToolToAgentApplier,
|
||||
)
|
||||
|
||||
applier = AddToolToAgentApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.TOOLS,
|
||||
op=EditOp.ADD_TOOL_TO_AGENT,
|
||||
target="agent.simple.tools",
|
||||
payload={"agent": "simple", "tool_name": "calculator"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
assert applier.validate(edit, ctx).ok
|
||||
|
||||
|
||||
class TestRemoveToolFromAgentApplier:
|
||||
"""Tests for RemoveToolFromAgentApplier."""
|
||||
|
||||
def test_apply_removes_tool(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
RemoveToolFromAgentApplier,
|
||||
)
|
||||
|
||||
applier = RemoveToolFromAgentApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.TOOLS,
|
||||
op=EditOp.REMOVE_TOOL_FROM_AGENT,
|
||||
target="agent.simple.tools",
|
||||
payload={"agent": "simple", "tool_name": "web_search"},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = ctx.config_path.read_text()
|
||||
assert "web_search" not in content
|
||||
|
||||
|
||||
class TestEditToolDescriptionApplier:
|
||||
"""Tests for EditToolDescriptionApplier."""
|
||||
|
||||
def test_apply_updates_description(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
EditToolDescriptionApplier,
|
||||
)
|
||||
|
||||
applier = EditToolDescriptionApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.TOOLS,
|
||||
op=EditOp.EDIT_TOOL_DESCRIPTION,
|
||||
target="tools.web_search.description",
|
||||
payload={
|
||||
"tool_name": "web_search",
|
||||
"new_description": "Search the internet for current information",
|
||||
},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = (ctx.tools_dir / "descriptions.toml").read_text()
|
||||
assert "Search the internet" in content
|
||||
|
||||
def test_apply_adds_new_tool_section(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.appliers.tools import (
|
||||
EditToolDescriptionApplier,
|
||||
)
|
||||
|
||||
applier = EditToolDescriptionApplier()
|
||||
ctx = _make_ctx(tmp_path)
|
||||
edit = Edit(
|
||||
id="edit-002",
|
||||
pillar=EditPillar.TOOLS,
|
||||
op=EditOp.EDIT_TOOL_DESCRIPTION,
|
||||
target="tools.calculator.description",
|
||||
payload={
|
||||
"tool_name": "calculator",
|
||||
"new_description": "Evaluate mathematical expressions",
|
||||
},
|
||||
rationale="test",
|
||||
expected_improvement="c1",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
applier.apply(edit, ctx)
|
||||
content = (ctx.tools_dir / "descriptions.toml").read_text()
|
||||
assert "[calculator]" in content
|
||||
assert "Evaluate mathematical" in content
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Tests for openjarvis.learning.distillation.gate.benchmark_gate module.
|
||||
|
||||
All tests use mock scorers — no live EvalRunner.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
|
||||
def _make_scorer(scores: dict[str, float], overall: float | None = None):
|
||||
"""Return a callable that produces a BenchmarkSnapshot with given scores."""
|
||||
|
||||
def scorer(
|
||||
*, benchmark_version: str, subsample_size: int, seed: int
|
||||
) -> BenchmarkSnapshot:
|
||||
computed = sum(scores.values()) / max(len(scores), 1)
|
||||
return BenchmarkSnapshot(
|
||||
benchmark_version=benchmark_version,
|
||||
overall_score=overall if overall is not None else computed,
|
||||
cluster_scores=scores,
|
||||
task_count=subsample_size,
|
||||
elapsed_seconds=5.0,
|
||||
)
|
||||
|
||||
return scorer
|
||||
|
||||
|
||||
class TestBenchmarkGate:
|
||||
"""Tests for BenchmarkGate."""
|
||||
|
||||
def test_accepts_improving_edit(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import (
|
||||
BenchmarkGate,
|
||||
)
|
||||
|
||||
before = BenchmarkSnapshot(
|
||||
benchmark_version="v1",
|
||||
overall_score=0.6,
|
||||
cluster_scores={"c1": 0.5, "c2": 0.7},
|
||||
task_count=50,
|
||||
elapsed_seconds=10.0,
|
||||
)
|
||||
gate = BenchmarkGate(
|
||||
scorer=_make_scorer({"c1": 0.6, "c2": 0.75}, overall=0.68),
|
||||
benchmark_version="v1",
|
||||
min_improvement=0.0,
|
||||
max_regression=0.05,
|
||||
subsample_size=50,
|
||||
)
|
||||
result = gate.evaluate(before=before, session_seed=42)
|
||||
assert result.accepted
|
||||
assert result.snapshot.overall_score == 0.68
|
||||
assert result.delta > 0
|
||||
|
||||
def test_rejects_no_improvement(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import (
|
||||
BenchmarkGate,
|
||||
)
|
||||
|
||||
before = BenchmarkSnapshot(
|
||||
benchmark_version="v1",
|
||||
overall_score=0.7,
|
||||
cluster_scores={"c1": 0.6, "c2": 0.8},
|
||||
task_count=50,
|
||||
elapsed_seconds=10.0,
|
||||
)
|
||||
gate = BenchmarkGate(
|
||||
scorer=_make_scorer({"c1": 0.6, "c2": 0.8}, overall=0.7),
|
||||
benchmark_version="v1",
|
||||
min_improvement=0.0,
|
||||
max_regression=0.05,
|
||||
subsample_size=50,
|
||||
)
|
||||
result = gate.evaluate(before=before, session_seed=42)
|
||||
assert not result.accepted
|
||||
assert "no improvement" in result.reason.lower()
|
||||
|
||||
def test_rejects_regression(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import (
|
||||
BenchmarkGate,
|
||||
)
|
||||
|
||||
before = BenchmarkSnapshot(
|
||||
benchmark_version="v1",
|
||||
overall_score=0.7,
|
||||
cluster_scores={"c1": 0.6, "c2": 0.8},
|
||||
task_count=50,
|
||||
elapsed_seconds=10.0,
|
||||
)
|
||||
# overall improves but c2 regresses badly
|
||||
gate = BenchmarkGate(
|
||||
scorer=_make_scorer({"c1": 0.75, "c2": 0.65}, overall=0.72),
|
||||
benchmark_version="v1",
|
||||
min_improvement=0.0,
|
||||
max_regression=0.05,
|
||||
subsample_size=50,
|
||||
)
|
||||
result = gate.evaluate(before=before, session_seed=42)
|
||||
assert not result.accepted
|
||||
assert "regression" in result.reason.lower()
|
||||
|
||||
def test_min_improvement_threshold(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import (
|
||||
BenchmarkGate,
|
||||
)
|
||||
|
||||
before = BenchmarkSnapshot(
|
||||
benchmark_version="v1",
|
||||
overall_score=0.7,
|
||||
cluster_scores={"c1": 0.6, "c2": 0.8},
|
||||
task_count=50,
|
||||
elapsed_seconds=10.0,
|
||||
)
|
||||
# Tiny improvement of 0.01, but min_improvement requires 0.05
|
||||
gate = BenchmarkGate(
|
||||
scorer=_make_scorer({"c1": 0.61, "c2": 0.81}, overall=0.71),
|
||||
benchmark_version="v1",
|
||||
min_improvement=0.05,
|
||||
max_regression=0.05,
|
||||
subsample_size=50,
|
||||
)
|
||||
result = gate.evaluate(before=before, session_seed=42)
|
||||
assert not result.accepted
|
||||
|
||||
def test_result_contains_snapshot(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.benchmark_gate import (
|
||||
BenchmarkGate,
|
||||
)
|
||||
|
||||
before = BenchmarkSnapshot(
|
||||
benchmark_version="v1",
|
||||
overall_score=0.5,
|
||||
cluster_scores={"c1": 0.5},
|
||||
task_count=50,
|
||||
elapsed_seconds=10.0,
|
||||
)
|
||||
gate = BenchmarkGate(
|
||||
scorer=_make_scorer({"c1": 0.7}, overall=0.7),
|
||||
benchmark_version="v1",
|
||||
min_improvement=0.0,
|
||||
max_regression=0.05,
|
||||
subsample_size=50,
|
||||
)
|
||||
result = gate.evaluate(before=before, session_seed=42)
|
||||
assert isinstance(result.snapshot, BenchmarkSnapshot)
|
||||
assert result.snapshot.benchmark_version == "v1"
|
||||
@@ -0,0 +1,218 @@
|
||||
"""Tests for openjarvis.learning.distillation.checkpoint.store module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _git(cwd: Path, *args: str) -> str:
|
||||
"""Helper to run git commands in tests."""
|
||||
result = subprocess.run(
|
||||
["git", *args],
|
||||
cwd=cwd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
)
|
||||
return result.stdout.strip()
|
||||
|
||||
|
||||
def _setup_isolated_repo_root(tmp_path: Path) -> Path:
|
||||
"""Create a fake openjarvis-home directory tree with config files for the
|
||||
CheckpointStore to track. Returns the root."""
|
||||
root = tmp_path / "openjarvis_home"
|
||||
(root / "agents" / "simple").mkdir(parents=True)
|
||||
(root / "tools").mkdir(parents=True)
|
||||
(root / "config.toml").write_text("[learning]\nenabled = true\n")
|
||||
(root / "agents" / "simple" / "system_prompt.md").write_text(
|
||||
"You are a helpful assistant.\n"
|
||||
)
|
||||
(root / "tools" / "descriptions.toml").write_text("[web_search]\n")
|
||||
return root
|
||||
|
||||
|
||||
class TestCheckpointStoreInit:
|
||||
"""Tests for CheckpointStore.init()."""
|
||||
|
||||
def test_creates_repo_with_baseline_commit(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
|
||||
assert (root / ".git").exists()
|
||||
log = _git(root, "log", "--oneline")
|
||||
assert "baseline" in log
|
||||
|
||||
def test_init_idempotent(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
first_sha = store.current_sha()
|
||||
|
||||
store.init() # Should not raise or create a second baseline.
|
||||
assert store.current_sha() == first_sha
|
||||
|
||||
def test_init_refuses_inside_source_tree(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
source_root = paths._find_source_root()
|
||||
assert source_root is not None
|
||||
bad_root = source_root / "fake_openjarvis_home"
|
||||
|
||||
store = CheckpointStore(bad_root)
|
||||
with pytest.raises(paths.ConfigurationError):
|
||||
store.init()
|
||||
|
||||
|
||||
class TestStageCommitDiscard:
|
||||
"""Tests for begin_stage / commit_stage / discard_stage."""
|
||||
|
||||
def test_commit_stage_creates_commit_with_trailers(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
|
||||
handle = store.begin_stage("edit-001")
|
||||
# Mutate a tracked file in the working tree.
|
||||
(root / "agents" / "simple" / "system_prompt.md").write_text(
|
||||
"You are a helpful, math-aware assistant.\n"
|
||||
)
|
||||
|
||||
new_sha = store.commit_stage(
|
||||
handle,
|
||||
message="learning: edit-001 add math hint",
|
||||
session_id="session-001",
|
||||
risk_tier="review",
|
||||
)
|
||||
|
||||
# New commit exists.
|
||||
assert new_sha != handle.pre_stage_sha
|
||||
# Commit message contains structured trailers.
|
||||
body = _git(root, "log", "-1", "--format=%B", new_sha)
|
||||
assert "Edit-ID: edit-001" in body
|
||||
assert "Session-ID: session-001" in body
|
||||
assert "Risk-Tier: review" in body
|
||||
|
||||
def test_discard_stage_restores_working_tree(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
original = (root / "agents" / "simple" / "system_prompt.md").read_text()
|
||||
|
||||
handle = store.begin_stage("edit-002")
|
||||
(root / "agents" / "simple" / "system_prompt.md").write_text(
|
||||
"Mutated content that should be discarded.\n"
|
||||
)
|
||||
|
||||
store.discard_stage(handle)
|
||||
|
||||
restored = (root / "agents" / "simple" / "system_prompt.md").read_text()
|
||||
assert restored == original
|
||||
# HEAD must equal the pre-stage sha.
|
||||
assert store.current_sha() == handle.pre_stage_sha
|
||||
|
||||
def test_begin_stage_refuses_dirty_working_tree(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
DirtyWorkingTreeError,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
|
||||
# Create an untracked, uncommitted change.
|
||||
(root / "agents" / "simple" / "system_prompt.md").write_text(
|
||||
"Pre-existing manual edit.\n"
|
||||
)
|
||||
|
||||
with pytest.raises(DirtyWorkingTreeError):
|
||||
store.begin_stage("edit-003")
|
||||
|
||||
|
||||
class TestRevertSession:
|
||||
"""Tests for revert_session."""
|
||||
|
||||
def test_revert_creates_new_commits_and_does_not_rewrite(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
|
||||
# Apply two commits tagged with the same session id.
|
||||
handle1 = store.begin_stage("edit-001")
|
||||
(root / "agents" / "simple" / "system_prompt.md").write_text("version A\n")
|
||||
store.commit_stage(
|
||||
handle1,
|
||||
message="learning: edit-001 v A",
|
||||
session_id="session-XYZ",
|
||||
risk_tier="auto",
|
||||
)
|
||||
|
||||
handle2 = store.begin_stage("edit-002")
|
||||
(root / "tools" / "descriptions.toml").write_text(
|
||||
"[web_search]\nupdated = true\n"
|
||||
)
|
||||
store.commit_stage(
|
||||
handle2,
|
||||
message="learning: edit-002 update tool",
|
||||
session_id="session-XYZ",
|
||||
risk_tier="auto",
|
||||
)
|
||||
|
||||
before_revert_log_count = len(_git(root, "log", "--oneline").splitlines())
|
||||
|
||||
revert_shas = store.revert_session("session-XYZ")
|
||||
assert len(revert_shas) == 2
|
||||
|
||||
# Two new commits added (the reverts), no history rewriting.
|
||||
after_revert_log_count = len(_git(root, "log", "--oneline").splitlines())
|
||||
assert after_revert_log_count == before_revert_log_count + 2
|
||||
|
||||
# Files restored to baseline.
|
||||
assert (
|
||||
root / "agents" / "simple" / "system_prompt.md"
|
||||
).read_text() == "You are a helpful assistant.\n"
|
||||
assert (root / "tools" / "descriptions.toml").read_text() == "[web_search]\n"
|
||||
|
||||
def test_revert_session_with_no_commits_returns_empty(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.checkpoint.store import (
|
||||
CheckpointStore,
|
||||
)
|
||||
|
||||
root = _setup_isolated_repo_root(tmp_path)
|
||||
store = CheckpointStore(root)
|
||||
store.init()
|
||||
|
||||
result = store.revert_session("session-with-no-commits")
|
||||
assert result == []
|
||||
@@ -0,0 +1,44 @@
|
||||
"""Tests for the jarvis learning CLI subcommand group."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from click.testing import CliRunner
|
||||
|
||||
|
||||
class TestLearningCLI:
|
||||
def test_learning_group_exists(self) -> None:
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
runner = CliRunner()
|
||||
result = runner.invoke(learning_group, ["--help"])
|
||||
assert result.exit_code == 0
|
||||
out = result.output.lower()
|
||||
assert "learning" in out or "distillation" in out
|
||||
|
||||
def test_init_subcommand(self) -> None:
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
runner = CliRunner()
|
||||
result = runner.invoke(learning_group, ["init", "--help"])
|
||||
assert result.exit_code == 0
|
||||
|
||||
def test_run_subcommand(self) -> None:
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
runner = CliRunner()
|
||||
result = runner.invoke(learning_group, ["run", "--help"])
|
||||
assert result.exit_code == 0
|
||||
|
||||
def test_history_subcommand(self) -> None:
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
runner = CliRunner()
|
||||
result = runner.invoke(learning_group, ["history", "--help"])
|
||||
assert result.exit_code == 0
|
||||
|
||||
def test_rollback_subcommand(self) -> None:
|
||||
from openjarvis.learning.distillation.cli import learning_group
|
||||
|
||||
runner = CliRunner()
|
||||
result = runner.invoke(learning_group, ["rollback", "--help"])
|
||||
assert result.exit_code == 0
|
||||
@@ -0,0 +1,101 @@
|
||||
"""Tests for openjarvis.learning.distillation.gate.cold_start module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
@dataclass
|
||||
class _StubTrace:
|
||||
trace_id: str = "t1"
|
||||
query: str = "test"
|
||||
feedback: Optional[float] = 0.8
|
||||
agent: str = "simple"
|
||||
model: str = "qwen"
|
||||
outcome: Optional[str] = "success"
|
||||
result: str = "answer"
|
||||
started_at: float = 1712534400.0
|
||||
ended_at: float = 1712534401.0
|
||||
steps: list = field(default_factory=list)
|
||||
messages: list = field(default_factory=list)
|
||||
total_tokens: int = 100
|
||||
total_latency_seconds: float = 1.0
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
engine: str = "ollama"
|
||||
|
||||
|
||||
def _make_trace_store(count: int = 0, high_feedback_count: int = 0) -> MagicMock:
|
||||
store = MagicMock()
|
||||
store.count.return_value = count
|
||||
# list_traces returns traces with high feedback
|
||||
high_traces = [
|
||||
_StubTrace(trace_id=f"t{i}", feedback=0.9) for i in range(high_feedback_count)
|
||||
]
|
||||
store.list_traces.return_value = high_traces
|
||||
return store
|
||||
|
||||
|
||||
class TestCheckReadiness:
|
||||
"""Tests for check_readiness()."""
|
||||
|
||||
def test_not_ready_with_no_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_readiness,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=0)
|
||||
result = check_readiness(store, min_traces=20)
|
||||
assert not result.ready
|
||||
assert "not enough traces" in result.message.lower()
|
||||
|
||||
def test_not_ready_with_few_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_readiness,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=10)
|
||||
result = check_readiness(store, min_traces=20)
|
||||
assert not result.ready
|
||||
|
||||
def test_ready_with_enough_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_readiness,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=25)
|
||||
result = check_readiness(store, min_traces=20)
|
||||
assert result.ready
|
||||
|
||||
|
||||
class TestCheckBenchmarkReady:
|
||||
"""Tests for check_benchmark_ready()."""
|
||||
|
||||
def test_not_ready_with_no_high_feedback_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_benchmark_ready,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=30, high_feedback_count=0)
|
||||
result = check_benchmark_ready(store, min_feedback=0.7, min_samples=10)
|
||||
assert not result.ready
|
||||
assert "benchmark" in result.message.lower()
|
||||
|
||||
def test_not_ready_with_few_high_feedback_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_benchmark_ready,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=30, high_feedback_count=5)
|
||||
result = check_benchmark_ready(store, min_feedback=0.7, min_samples=10)
|
||||
assert not result.ready
|
||||
|
||||
def test_ready_with_enough_high_feedback_traces(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_benchmark_ready,
|
||||
)
|
||||
|
||||
store = _make_trace_store(count=30, high_feedback_count=15)
|
||||
result = check_benchmark_ready(store, min_feedback=0.7, min_samples=10)
|
||||
assert result.ready
|
||||
@@ -0,0 +1,282 @@
|
||||
"""Tests for openjarvis.learning.distillation.diagnose.tools module.
|
||||
|
||||
All tests use fixture stubs — no live TraceStore, CloudEngine, or ToolRegistry.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stubs for dependencies
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class _StubTrace:
|
||||
"""Minimal stub matching the Trace fields that tools access."""
|
||||
|
||||
trace_id: str = "trace-001"
|
||||
query: str = "What is 2+2?"
|
||||
agent: str = "simple"
|
||||
model: str = "qwen2.5-coder:7b"
|
||||
outcome: Optional[str] = "success"
|
||||
feedback: Optional[float] = 0.8
|
||||
started_at: float = 1712534400.0
|
||||
result: str = "4"
|
||||
steps: list = field(default_factory=list)
|
||||
messages: list = field(default_factory=list)
|
||||
total_tokens: int = 100
|
||||
total_latency_seconds: float = 1.0
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
ended_at: float = 1712534401.0
|
||||
engine: str = "ollama"
|
||||
|
||||
|
||||
def _make_stub_trace_store(traces: list[_StubTrace] | None = None) -> MagicMock:
|
||||
"""Create a mock TraceStore with canned responses."""
|
||||
store = MagicMock()
|
||||
traces = traces or [_StubTrace()]
|
||||
store.list_traces.return_value = traces
|
||||
store.get.side_effect = lambda tid: next(
|
||||
(t for t in traces if t.trace_id == tid), None
|
||||
)
|
||||
store.search.return_value = [
|
||||
{"trace_id": t.trace_id, "query": t.query, "score": 1.0} for t in traces
|
||||
]
|
||||
return store
|
||||
|
||||
|
||||
def _make_stub_benchmark_samples() -> list:
|
||||
"""Return a list of stub PersonalBenchmarkSample objects."""
|
||||
sample = MagicMock()
|
||||
sample.trace_id = "task-001"
|
||||
sample.query = "What is quantum computing?"
|
||||
sample.reference_answer = "Quantum computing uses qubits..."
|
||||
sample.category = "reasoning"
|
||||
sample.feedback_score = 0.9
|
||||
return [sample]
|
||||
|
||||
|
||||
def _make_stub_config(tmp_path: Path) -> dict:
|
||||
"""Create a minimal config dict and on-disk files."""
|
||||
agents_dir = tmp_path / "agents" / "simple"
|
||||
agents_dir.mkdir(parents=True)
|
||||
(agents_dir / "system_prompt.md").write_text("You are a helpful assistant.\n")
|
||||
tools_dir = tmp_path / "tools"
|
||||
tools_dir.mkdir(parents=True)
|
||||
(tools_dir / "descriptions.toml").write_text(
|
||||
'[web_search]\ndescription = "Search the web"\n'
|
||||
)
|
||||
config_path = tmp_path / "config.toml"
|
||||
config_path.write_text("[learning]\nenabled = true\n")
|
||||
return {
|
||||
"config_path": config_path,
|
||||
"openjarvis_home": tmp_path,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildDiagnosticTools:
|
||||
"""Tests for the build_diagnostic_tools factory."""
|
||||
|
||||
def test_returns_expected_tool_names(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=_make_stub_benchmark_samples(),
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
names = {t.name for t in tools}
|
||||
assert "list_traces" in names
|
||||
assert "get_trace" in names
|
||||
assert "search_traces" in names
|
||||
assert "get_current_config" in names
|
||||
assert "get_agent_prompt" in names
|
||||
assert "get_tool_description" in names
|
||||
assert "list_available_tools" in names
|
||||
assert "list_personal_benchmark" in names
|
||||
assert "run_student_on_task" in names
|
||||
assert "run_self_on_task" in names
|
||||
assert "compare_outputs" in names
|
||||
|
||||
def test_all_tools_have_openai_format(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=_make_stub_benchmark_samples(),
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
for tool in tools:
|
||||
spec = tool.to_openai_function()
|
||||
assert spec["type"] == "function"
|
||||
assert "name" in spec["function"]
|
||||
assert "parameters" in spec["function"]
|
||||
|
||||
|
||||
class TestListTraces:
|
||||
"""Tests for the list_traces diagnostic tool."""
|
||||
|
||||
def test_returns_trace_metas(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
traces = [
|
||||
_StubTrace(trace_id="t1", feedback=0.3),
|
||||
_StubTrace(trace_id="t2", feedback=0.9),
|
||||
]
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(traces),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
list_traces = next(t for t in tools if t.name == "list_traces")
|
||||
result = list_traces.fn(limit=10)
|
||||
parsed = json.loads(result)
|
||||
assert len(parsed) == 2
|
||||
assert parsed[0]["trace_id"] == "t1"
|
||||
|
||||
|
||||
class TestGetTrace:
|
||||
"""Tests for the get_trace diagnostic tool."""
|
||||
|
||||
def test_returns_trace_details(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
get_trace = next(t for t in tools if t.name == "get_trace")
|
||||
result = get_trace.fn(trace_id="trace-001")
|
||||
parsed = json.loads(result)
|
||||
assert parsed["trace_id"] == "trace-001"
|
||||
assert parsed["query"] == "What is 2+2?"
|
||||
|
||||
def test_returns_error_for_unknown_trace(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
get_trace = next(t for t in tools if t.name == "get_trace")
|
||||
result = get_trace.fn(trace_id="nonexistent")
|
||||
assert "not found" in result.lower()
|
||||
|
||||
|
||||
class TestGetCurrentConfig:
|
||||
"""Tests for the get_current_config diagnostic tool."""
|
||||
|
||||
def test_returns_config_content(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
get_config = next(t for t in tools if t.name == "get_current_config")
|
||||
result = get_config.fn()
|
||||
assert "learning" in result
|
||||
|
||||
|
||||
class TestGetAgentPrompt:
|
||||
"""Tests for the get_agent_prompt diagnostic tool."""
|
||||
|
||||
def test_returns_prompt_content(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
get_prompt = next(t for t in tools if t.name == "get_agent_prompt")
|
||||
result = get_prompt.fn(agent_name="simple")
|
||||
assert "helpful assistant" in result
|
||||
|
||||
|
||||
class TestListPersonalBenchmark:
|
||||
"""Tests for the list_personal_benchmark diagnostic tool."""
|
||||
|
||||
def test_returns_benchmark_tasks(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.tools import (
|
||||
build_diagnostic_tools,
|
||||
)
|
||||
|
||||
tools = build_diagnostic_tools(
|
||||
trace_store=_make_stub_trace_store(),
|
||||
config=_make_stub_config(tmp_path),
|
||||
benchmark_samples=_make_stub_benchmark_samples(),
|
||||
student_runner=MagicMock(),
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
judge=MagicMock(),
|
||||
session_id="session-001",
|
||||
)
|
||||
list_bench = next(t for t in tools if t.name == "list_personal_benchmark")
|
||||
result = list_bench.fn(limit=10)
|
||||
parsed = json.loads(result)
|
||||
assert len(parsed) == 1
|
||||
assert parsed[0]["task_id"] == "task-001"
|
||||
@@ -0,0 +1,173 @@
|
||||
"""Tests for openjarvis.learning.distillation.diagnose.types module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
|
||||
class TestTraceMeta:
|
||||
"""Tests for TraceMeta dataclass."""
|
||||
|
||||
def test_constructs_with_required_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import TraceMeta
|
||||
|
||||
meta = TraceMeta(
|
||||
trace_id="trace-001",
|
||||
query="What is 2+2?",
|
||||
agent="simple",
|
||||
model="qwen2.5-coder:7b",
|
||||
outcome="success",
|
||||
feedback=0.8,
|
||||
started_at=1712534400.0,
|
||||
)
|
||||
assert meta.trace_id == "trace-001"
|
||||
assert meta.feedback == 0.8
|
||||
|
||||
def test_feedback_can_be_none(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import TraceMeta
|
||||
|
||||
meta = TraceMeta(
|
||||
trace_id="trace-002",
|
||||
query="test",
|
||||
agent="simple",
|
||||
model="qwen2.5-coder:7b",
|
||||
outcome=None,
|
||||
feedback=None,
|
||||
started_at=1712534400.0,
|
||||
)
|
||||
assert meta.feedback is None
|
||||
|
||||
|
||||
class TestBenchmarkTask:
|
||||
"""Tests for BenchmarkTask dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import BenchmarkTask
|
||||
|
||||
task = BenchmarkTask(
|
||||
task_id="task-001",
|
||||
query="Explain quantum computing",
|
||||
reference_answer="Quantum computing uses qubits...",
|
||||
category="reasoning",
|
||||
)
|
||||
assert task.task_id == "task-001"
|
||||
assert task.category == "reasoning"
|
||||
|
||||
|
||||
class TestStudentRun:
|
||||
"""Tests for StudentRun dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import StudentRun
|
||||
|
||||
run = StudentRun(
|
||||
task_id="task-001",
|
||||
output="The answer is 4.",
|
||||
score=0.9,
|
||||
trace_id="trace-new-001",
|
||||
latency_seconds=2.5,
|
||||
tokens_used=150,
|
||||
)
|
||||
assert run.score == 0.9
|
||||
assert run.trace_id == "trace-new-001"
|
||||
|
||||
|
||||
class TestTeacherRun:
|
||||
"""Tests for TeacherRun dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import TeacherRun
|
||||
|
||||
run = TeacherRun(
|
||||
task_id="task-001",
|
||||
output="Quantum computing is...",
|
||||
reasoning="I approached this by...",
|
||||
cost_usd=0.05,
|
||||
tokens_used=500,
|
||||
)
|
||||
assert run.cost_usd == 0.05
|
||||
|
||||
|
||||
class TestComparisonResult:
|
||||
"""Tests for ComparisonResult dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import ComparisonResult
|
||||
|
||||
result = ComparisonResult(
|
||||
task_id="task-001",
|
||||
student_score=0.3,
|
||||
teacher_score=0.9,
|
||||
judge_reasoning="The student missed the key concept...",
|
||||
)
|
||||
assert result.student_score == 0.3
|
||||
assert result.teacher_score == 0.9
|
||||
|
||||
|
||||
class TestToolMeta:
|
||||
"""Tests for ToolMeta dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import ToolMeta
|
||||
|
||||
meta = ToolMeta(
|
||||
name="calculator",
|
||||
description="Evaluate math expressions",
|
||||
category="math",
|
||||
agents=["simple", "react"],
|
||||
)
|
||||
assert meta.name == "calculator"
|
||||
assert len(meta.agents) == 2
|
||||
|
||||
|
||||
class TestDiagnosticTool:
|
||||
"""Tests for DiagnosticTool dataclass."""
|
||||
|
||||
def test_constructs_with_callable(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import DiagnosticTool
|
||||
|
||||
def my_func(**kwargs: object) -> str:
|
||||
return "result"
|
||||
|
||||
tool = DiagnosticTool(
|
||||
name="test_tool",
|
||||
description="A test tool",
|
||||
parameters={"type": "object", "properties": {}},
|
||||
fn=my_func,
|
||||
)
|
||||
assert tool.name == "test_tool"
|
||||
assert tool.fn(foo="bar") == "result"
|
||||
|
||||
|
||||
class TestToolCallRecord:
|
||||
"""Tests for ToolCallRecord dataclass."""
|
||||
|
||||
def test_constructs(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import ToolCallRecord
|
||||
|
||||
record = ToolCallRecord(
|
||||
timestamp=datetime(2026, 4, 9, 3, 0, 0, tzinfo=timezone.utc),
|
||||
tool="list_traces",
|
||||
args={"limit": 10},
|
||||
result="[...]",
|
||||
latency_ms=42.5,
|
||||
cost_usd=0.0,
|
||||
)
|
||||
assert record.tool == "list_traces"
|
||||
assert record.latency_ms == 42.5
|
||||
|
||||
def test_to_jsonl_dict(self) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.types import ToolCallRecord
|
||||
|
||||
record = ToolCallRecord(
|
||||
timestamp=datetime(2026, 4, 9, 3, 0, 0, tzinfo=timezone.utc),
|
||||
tool="get_trace",
|
||||
args={"trace_id": "t1"},
|
||||
result="trace data",
|
||||
latency_ms=10.0,
|
||||
cost_usd=0.01,
|
||||
)
|
||||
d = record.to_jsonl_dict()
|
||||
assert d["tool"] == "get_trace"
|
||||
assert d["timestamp"] == "2026-04-09T03:00:00+00:00"
|
||||
assert d["cost_usd"] == 0.01
|
||||
@@ -0,0 +1,235 @@
|
||||
"""Tests for openjarvis.learning.distillation.diagnose.runner module.
|
||||
|
||||
All tests use mocked dependencies — no live API calls.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def _make_engine_response(content: str, tool_calls: list | None = None) -> dict:
|
||||
"""Create a mock engine.generate() response."""
|
||||
resp = {
|
||||
"content": content,
|
||||
"finish_reason": "stop" if not tool_calls else "tool_calls",
|
||||
"usage": {"prompt_tokens": 100, "completion_tokens": 200, "total_tokens": 300},
|
||||
"cost_usd": 0.02,
|
||||
}
|
||||
if tool_calls:
|
||||
resp["tool_calls"] = tool_calls
|
||||
return resp
|
||||
|
||||
|
||||
def _make_diagnosis_content() -> str:
|
||||
"""A plausible teacher diagnosis output with embedded cluster JSON."""
|
||||
return (
|
||||
"## Diagnosis\n\n"
|
||||
"The student has two main failure patterns:\n\n"
|
||||
"### Cluster 1: Math Routing\n"
|
||||
"Math queries are being routed to qwen-3b which lacks chain-of-thought.\n\n"
|
||||
"### Cluster 2: Tool Selection\n"
|
||||
"The student frequently fails to use the calculator tool for arithmetic.\n\n"
|
||||
"```json\n"
|
||||
+ json.dumps(
|
||||
[
|
||||
{
|
||||
"id": "cluster-001",
|
||||
"description": "Math queries routed to qwen-3b",
|
||||
"sample_trace_ids": ["t1", "t2", "t3"],
|
||||
"student_failure_rate": 0.8,
|
||||
"teacher_success_rate": 0.95,
|
||||
"skill_gap": "Student lacks chain-of-thought on multi-step math",
|
||||
},
|
||||
{
|
||||
"id": "cluster-002",
|
||||
"description": "Calculator tool not used for arithmetic",
|
||||
"sample_trace_ids": ["t4", "t5", "t6"],
|
||||
"student_failure_rate": 0.6,
|
||||
"teacher_success_rate": 0.9,
|
||||
"skill_gap": "Student does not invoke calculator tool",
|
||||
},
|
||||
]
|
||||
)
|
||||
+ "\n```\n"
|
||||
)
|
||||
|
||||
|
||||
class TestDiagnosisRunner:
|
||||
"""Tests for DiagnosisRunner."""
|
||||
|
||||
def test_produces_diagnosis_artifact(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_engine_response(
|
||||
content=_make_diagnosis_content()
|
||||
)
|
||||
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_dir=tmp_path / "session-001",
|
||||
session_id="session-001",
|
||||
config={
|
||||
"config_path": tmp_path / "config.toml",
|
||||
"openjarvis_home": tmp_path,
|
||||
},
|
||||
)
|
||||
# Create minimal config file
|
||||
(tmp_path / "config.toml").write_text("[learning]\n")
|
||||
|
||||
runner.run()
|
||||
|
||||
# Diagnosis artifact written
|
||||
diagnosis_path = tmp_path / "session-001" / "diagnosis.md"
|
||||
assert diagnosis_path.exists()
|
||||
assert "Math" in diagnosis_path.read_text()
|
||||
|
||||
def test_returns_failure_clusters(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_engine_response(
|
||||
content=_make_diagnosis_content()
|
||||
)
|
||||
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_dir=tmp_path / "session-001",
|
||||
session_id="session-001",
|
||||
config={
|
||||
"config_path": tmp_path / "config.toml",
|
||||
"openjarvis_home": tmp_path,
|
||||
},
|
||||
)
|
||||
(tmp_path / "config.toml").write_text("[learning]\n")
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert len(result.clusters) == 2
|
||||
assert result.clusters[0].id == "cluster-001"
|
||||
assert result.clusters[0].student_failure_rate == 0.8
|
||||
assert result.clusters[1].id == "cluster-002"
|
||||
|
||||
def test_persists_teacher_traces_jsonl(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
engine = MagicMock()
|
||||
# Turn 1: tool call
|
||||
engine.generate.side_effect = [
|
||||
_make_engine_response(
|
||||
content="",
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "c1",
|
||||
"name": "get_current_config",
|
||||
"arguments": "{}",
|
||||
}
|
||||
],
|
||||
),
|
||||
_make_engine_response(content=_make_diagnosis_content()),
|
||||
]
|
||||
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_dir=tmp_path / "session-001",
|
||||
session_id="session-001",
|
||||
config={
|
||||
"config_path": tmp_path / "config.toml",
|
||||
"openjarvis_home": tmp_path,
|
||||
},
|
||||
)
|
||||
(tmp_path / "config.toml").write_text("[learning]\n")
|
||||
|
||||
runner.run()
|
||||
|
||||
jsonl_path = tmp_path / "session-001" / "teacher_traces" / "diagnose.jsonl"
|
||||
assert jsonl_path.exists()
|
||||
lines = jsonl_path.read_text().strip().splitlines()
|
||||
assert len(lines) >= 1
|
||||
record = json.loads(lines[0])
|
||||
assert "tool" in record
|
||||
|
||||
def test_returns_cost(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_engine_response(
|
||||
content=_make_diagnosis_content()
|
||||
)
|
||||
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_dir=tmp_path / "session-001",
|
||||
session_id="session-001",
|
||||
config={
|
||||
"config_path": tmp_path / "config.toml",
|
||||
"openjarvis_home": tmp_path,
|
||||
},
|
||||
)
|
||||
(tmp_path / "config.toml").write_text("[learning]\n")
|
||||
|
||||
result = runner.run()
|
||||
assert result.cost_usd >= 0.0
|
||||
|
||||
def test_handles_no_clusters_in_output(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_engine_response(
|
||||
content="## Diagnosis\nNo clear failure patterns found."
|
||||
)
|
||||
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_dir=tmp_path / "session-001",
|
||||
session_id="session-001",
|
||||
config={
|
||||
"config_path": tmp_path / "config.toml",
|
||||
"openjarvis_home": tmp_path,
|
||||
},
|
||||
)
|
||||
(tmp_path / "config.toml").write_text("[learning]\n")
|
||||
|
||||
result = runner.run()
|
||||
|
||||
assert result.clusters == []
|
||||
# Diagnosis artifact is still written
|
||||
assert (tmp_path / "session-001" / "diagnosis.md").exists()
|
||||
@@ -0,0 +1,150 @@
|
||||
"""Tests for the per-edit execution loop."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.execute.base import ApplyContext
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_ctx(tmp_path: Path) -> ApplyContext:
|
||||
(tmp_path / "config.toml").write_text(
|
||||
'[learning.routing.policy_map]\nmath = "qwen2.5-coder:3b"\n'
|
||||
)
|
||||
agents_dir = tmp_path / "agents" / "simple"
|
||||
agents_dir.mkdir(parents=True)
|
||||
(agents_dir / "system_prompt.md").write_text("You are helpful.\n")
|
||||
tools_dir = tmp_path / "tools"
|
||||
tools_dir.mkdir(parents=True)
|
||||
(tools_dir / "descriptions.toml").write_text(
|
||||
'[web_search]\ndescription = "Search"\n'
|
||||
)
|
||||
return ApplyContext(openjarvis_home=tmp_path, session_id="s1")
|
||||
|
||||
|
||||
def _make_auto_edit(edit_id: str = "edit-001") -> Edit:
|
||||
return Edit(
|
||||
id=edit_id,
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
target="routing.math",
|
||||
payload={"query_class": "math", "model": "qwen2.5-coder:14b"},
|
||||
rationale="Route math to bigger model",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
|
||||
|
||||
def _make_review_edit(edit_id: str = "edit-002") -> Edit:
|
||||
return Edit(
|
||||
id=edit_id,
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"new_content": "New prompt.\n"},
|
||||
rationale="Better prompt",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
|
||||
|
||||
def _make_lora_edit(edit_id: str = "edit-lora") -> Edit:
|
||||
return Edit(
|
||||
id=edit_id,
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.LORA_FINETUNE,
|
||||
target="models.qwen",
|
||||
payload={"target_model": "qwen", "data_source": "all"},
|
||||
rationale="Fine tune",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.MANUAL,
|
||||
)
|
||||
|
||||
|
||||
class TestExecuteEdits:
|
||||
"""Tests for execute_edits()."""
|
||||
|
||||
def test_applies_auto_tier_edit(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[_make_auto_edit()],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.TIERED,
|
||||
)
|
||||
assert len(outcomes) == 1
|
||||
assert outcomes[0].status == "applied"
|
||||
|
||||
def test_review_edit_goes_to_pending_in_tiered_mode(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[_make_review_edit()],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.TIERED,
|
||||
)
|
||||
assert len(outcomes) == 1
|
||||
assert outcomes[0].status == "pending_review"
|
||||
|
||||
def test_review_edit_applied_in_auto_mode(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[_make_review_edit()],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.AUTO,
|
||||
)
|
||||
assert len(outcomes) == 1
|
||||
assert outcomes[0].status == "applied"
|
||||
|
||||
def test_manual_tier_skipped(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[_make_lora_edit()],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.TIERED,
|
||||
)
|
||||
assert len(outcomes) == 1
|
||||
assert outcomes[0].status == "skipped"
|
||||
|
||||
def test_all_edits_pending_in_manual_mode(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[_make_auto_edit()],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.MANUAL,
|
||||
)
|
||||
assert len(outcomes) == 1
|
||||
assert outcomes[0].status == "pending_review"
|
||||
|
||||
def test_multiple_edits_processed(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.execute.loop import execute_edits
|
||||
|
||||
ctx = _make_ctx(tmp_path)
|
||||
outcomes = execute_edits(
|
||||
edits=[
|
||||
_make_auto_edit("e1"),
|
||||
_make_review_edit("e2"),
|
||||
_make_lora_edit("e3"),
|
||||
],
|
||||
ctx=ctx,
|
||||
autonomy_mode=AutonomyMode.TIERED,
|
||||
)
|
||||
assert len(outcomes) == 3
|
||||
assert outcomes[0].status == "applied"
|
||||
assert outcomes[1].status == "pending_review"
|
||||
assert outcomes[2].status == "skipped"
|
||||
@@ -0,0 +1,177 @@
|
||||
"""Live integration tests for the distillation subsystem.
|
||||
|
||||
These tests use REAL API calls (CloudEngine with Anthropic) and real
|
||||
TraceStore data. They are gated on the ``cloud`` marker — skip them
|
||||
with ``pytest -m "not cloud"``.
|
||||
|
||||
Requires:
|
||||
- ANTHROPIC_API_KEY environment variable set
|
||||
- TraceStore at ~/.openjarvis/traces.db with some traces
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# Skip entire module if no API key
|
||||
pytestmark = pytest.mark.cloud
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def anthropic_key():
|
||||
key = os.environ.get("ANTHROPIC_API_KEY")
|
||||
if not key:
|
||||
pytest.skip("ANTHROPIC_API_KEY not set")
|
||||
return key
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def cloud_engine(anthropic_key):
|
||||
from openjarvis.engine.cloud import CloudEngine
|
||||
|
||||
return CloudEngine()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def real_trace_store():
|
||||
from openjarvis.traces.store import TraceStore
|
||||
|
||||
db_path = Path.home() / ".openjarvis" / "traces.db"
|
||||
if not db_path.exists():
|
||||
pytest.skip("No traces.db found at ~/.openjarvis/")
|
||||
store = TraceStore(db_path)
|
||||
if store.count() < 5:
|
||||
pytest.skip("Need at least 5 traces for live test")
|
||||
return store
|
||||
|
||||
|
||||
class TestCloudEngineDirectCall:
|
||||
"""Verify CloudEngine works with real API."""
|
||||
|
||||
def test_generate_produces_content(self, cloud_engine) -> None:
|
||||
from openjarvis.core.types import Message, Role
|
||||
|
||||
result = cloud_engine.generate(
|
||||
messages=[
|
||||
Message(
|
||||
role=Role.USER,
|
||||
content="What is 2+2? Answer with just the number.",
|
||||
)
|
||||
],
|
||||
model="claude-sonnet-4-6",
|
||||
max_tokens=10,
|
||||
)
|
||||
assert "content" in result
|
||||
assert "4" in result["content"]
|
||||
assert result.get("cost_usd", 0) > 0
|
||||
|
||||
|
||||
class TestTeacherAgentLive:
|
||||
"""Test TeacherAgent with a real CloudEngine."""
|
||||
|
||||
def test_teacher_agent_single_turn(self, cloud_engine) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.teacher_agent import (
|
||||
TeacherAgent,
|
||||
)
|
||||
|
||||
agent = TeacherAgent(
|
||||
engine=cloud_engine,
|
||||
model="claude-sonnet-4-6",
|
||||
tools=[],
|
||||
max_turns=2,
|
||||
max_cost_usd=0.50,
|
||||
)
|
||||
result = agent.run(
|
||||
"You are being tested. Simply respond with: 'TeacherAgent works.'",
|
||||
system_prompt="You are a test assistant. Follow instructions exactly.",
|
||||
)
|
||||
assert result.content
|
||||
assert result.turns >= 1
|
||||
assert result.total_cost_usd > 0
|
||||
print(f" Teacher response: {result.content[:100]}")
|
||||
print(f" Cost: ${result.total_cost_usd:.4f}, Turns: {result.turns}")
|
||||
|
||||
|
||||
class TestDiagnosisRunnerLive:
|
||||
"""Test DiagnosisRunner with real CloudEngine + real traces."""
|
||||
|
||||
def test_diagnosis_produces_output(
|
||||
self, cloud_engine, real_trace_store, tmp_path
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.diagnose.runner import (
|
||||
DiagnosisRunner,
|
||||
)
|
||||
|
||||
# Create minimal config
|
||||
config_dir = tmp_path / "oj_home"
|
||||
config_dir.mkdir()
|
||||
(config_dir / "config.toml").write_text("[learning]\nenabled = true\n")
|
||||
|
||||
session_dir = tmp_path / "session"
|
||||
runner = DiagnosisRunner(
|
||||
teacher_engine=cloud_engine,
|
||||
teacher_model="claude-sonnet-4-6",
|
||||
trace_store=real_trace_store,
|
||||
benchmark_samples=[],
|
||||
student_runner=lambda q, **kw: type(
|
||||
"R", (), {"content": "mock", "score": 0.5}
|
||||
)(),
|
||||
judge=type("J", (), {"score_trace": lambda self, t: (0.5, "mock")})(),
|
||||
session_dir=session_dir,
|
||||
session_id="live-test-001",
|
||||
config={
|
||||
"config_path": config_dir / "config.toml",
|
||||
"openjarvis_home": config_dir,
|
||||
},
|
||||
max_turns=5, # Keep it cheap
|
||||
max_cost_usd=1.0,
|
||||
)
|
||||
|
||||
result = runner.run()
|
||||
|
||||
# Diagnosis should produce output
|
||||
assert result.diagnosis_md, "Diagnosis produced no markdown"
|
||||
assert len(result.diagnosis_md) > 50, "Diagnosis too short"
|
||||
print(f" Diagnosis length: {len(result.diagnosis_md)} chars")
|
||||
print(f" Clusters found: {len(result.clusters)}")
|
||||
print(f" Cost: ${result.cost_usd:.4f}")
|
||||
print(f" Tool calls: {len(result.tool_call_records)}")
|
||||
|
||||
# Artifacts should exist
|
||||
assert (session_dir / "diagnosis.md").exists()
|
||||
assert (session_dir / "teacher_traces" / "diagnose.jsonl").exists()
|
||||
|
||||
# Print first 200 chars of diagnosis
|
||||
print(f" Diagnosis preview: {result.diagnosis_md[:200]}...")
|
||||
|
||||
|
||||
class TestColdStartLive:
|
||||
"""Test cold start behavior with real trace store."""
|
||||
|
||||
def test_orchestrator_cold_start_with_no_feedback(
|
||||
self, cloud_engine, real_trace_store, tmp_path
|
||||
) -> None:
|
||||
"""With 373 traces but 0 feedback, the orchestrator should handle
|
||||
this gracefully — either by running (traces > 20) or by giving
|
||||
a clear message about what's missing."""
|
||||
from openjarvis.learning.distillation.gate.cold_start import (
|
||||
check_benchmark_ready,
|
||||
check_readiness,
|
||||
)
|
||||
|
||||
# Check trace readiness
|
||||
trace_ready = check_readiness(real_trace_store, min_traces=20)
|
||||
print(f" Trace readiness: {trace_ready.ready} ({trace_ready.message})")
|
||||
|
||||
# Check benchmark readiness
|
||||
bench_ready = check_benchmark_ready(
|
||||
real_trace_store, min_feedback=0.7, min_samples=10
|
||||
)
|
||||
print(f" Benchmark readiness: {bench_ready.ready} ({bench_ready.message})")
|
||||
|
||||
# With 373 traces but 0 feedback: traces ready, benchmark not ready
|
||||
assert trace_ready.ready, "Should have enough traces"
|
||||
assert not bench_ready.ready, "Should not have enough high-feedback traces"
|
||||
@@ -0,0 +1,586 @@
|
||||
"""Tests for openjarvis.learning.distillation.models module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Enums
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEditPillar:
|
||||
"""Tests for EditPillar enum."""
|
||||
|
||||
def test_has_four_pillars(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditPillar
|
||||
|
||||
assert EditPillar.INTELLIGENCE.value == "intelligence"
|
||||
assert EditPillar.AGENT.value == "agent"
|
||||
assert EditPillar.TOOLS.value == "tools"
|
||||
assert EditPillar.ENGINE.value == "engine"
|
||||
|
||||
def test_is_string_enum(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditPillar
|
||||
|
||||
assert isinstance(EditPillar.AGENT, str)
|
||||
assert EditPillar("agent") is EditPillar.AGENT
|
||||
|
||||
|
||||
class TestEditRiskTier:
|
||||
"""Tests for EditRiskTier enum."""
|
||||
|
||||
def test_has_three_tiers(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditRiskTier
|
||||
|
||||
assert EditRiskTier.AUTO.value == "auto"
|
||||
assert EditRiskTier.REVIEW.value == "review"
|
||||
assert EditRiskTier.MANUAL.value == "manual"
|
||||
|
||||
|
||||
class TestEditOp:
|
||||
"""Tests for EditOp enum — must contain all v1 ops plus v2 placeholders."""
|
||||
|
||||
def test_intelligence_ops(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
assert EditOp.SET_MODEL_FOR_QUERY_CLASS.value == "set_model_for_query_class"
|
||||
assert EditOp.SET_MODEL_PARAM.value == "set_model_param"
|
||||
|
||||
def test_agent_ops(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
assert EditOp.PATCH_SYSTEM_PROMPT.value == "patch_system_prompt"
|
||||
assert EditOp.REPLACE_SYSTEM_PROMPT.value == "replace_system_prompt"
|
||||
assert EditOp.SET_AGENT_CLASS.value == "set_agent_class"
|
||||
assert EditOp.SET_AGENT_PARAM.value == "set_agent_param"
|
||||
assert EditOp.EDIT_FEW_SHOT_EXEMPLARS.value == "edit_few_shot_exemplars"
|
||||
|
||||
def test_tools_ops(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
assert EditOp.ADD_TOOL_TO_AGENT.value == "add_tool_to_agent"
|
||||
assert EditOp.REMOVE_TOOL_FROM_AGENT.value == "remove_tool_from_agent"
|
||||
assert EditOp.EDIT_TOOL_DESCRIPTION.value == "edit_tool_description"
|
||||
|
||||
def test_v2_placeholder_ops(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
|
||||
assert EditOp.LORA_FINETUNE.value == "lora_finetune"
|
||||
|
||||
|
||||
class TestTriggerKind:
|
||||
"""Tests for TriggerKind enum."""
|
||||
|
||||
def test_four_trigger_kinds(self) -> None:
|
||||
from openjarvis.learning.distillation.models import TriggerKind
|
||||
|
||||
assert TriggerKind.SCHEDULED.value == "scheduled"
|
||||
assert TriggerKind.CLUSTER.value == "cluster"
|
||||
assert TriggerKind.USER_FLAG.value == "user_flag"
|
||||
assert TriggerKind.ON_DEMAND.value == "on_demand"
|
||||
|
||||
|
||||
class TestAutonomyMode:
|
||||
"""Tests for AutonomyMode enum."""
|
||||
|
||||
def test_three_modes(self) -> None:
|
||||
from openjarvis.learning.distillation.models import AutonomyMode
|
||||
|
||||
assert AutonomyMode.AUTO.value == "auto"
|
||||
assert AutonomyMode.TIERED.value == "tiered"
|
||||
assert AutonomyMode.MANUAL.value == "manual"
|
||||
|
||||
|
||||
class TestSessionStatus:
|
||||
"""Tests for SessionStatus enum."""
|
||||
|
||||
def test_all_statuses(self) -> None:
|
||||
from openjarvis.learning.distillation.models import SessionStatus
|
||||
|
||||
assert SessionStatus.INITIATED.value == "initiated"
|
||||
assert SessionStatus.DIAGNOSING.value == "diagnosing"
|
||||
assert SessionStatus.PLANNING.value == "planning"
|
||||
assert SessionStatus.EXECUTING.value == "executing"
|
||||
assert SessionStatus.AWAITING_REVIEW.value == "awaiting_review"
|
||||
assert SessionStatus.COMPLETED.value == "completed"
|
||||
assert SessionStatus.FAILED.value == "failed"
|
||||
assert SessionStatus.ROLLED_BACK.value == "rolled_back"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Edit
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEdit:
|
||||
"""Tests for Edit pydantic model."""
|
||||
|
||||
def _valid_edit_kwargs(self) -> dict:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
return {
|
||||
"id": "11111111-2222-3333-4444-555555555555",
|
||||
"pillar": EditPillar.INTELLIGENCE,
|
||||
"op": EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
"target": "learning.routing.policy_map.math",
|
||||
"payload": {"query_class": "math", "model": "qwen2.5-coder:14b"},
|
||||
"rationale": "Math queries are misrouted to qwen-3b",
|
||||
"expected_improvement": "math_failures cluster",
|
||||
"risk_tier": EditRiskTier.AUTO,
|
||||
"references": ["trace-001", "trace-002"],
|
||||
}
|
||||
|
||||
def test_constructs_with_valid_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
edit = Edit(**self._valid_edit_kwargs())
|
||||
|
||||
assert edit.id == "11111111-2222-3333-4444-555555555555"
|
||||
assert edit.target == "learning.routing.policy_map.math"
|
||||
assert edit.payload == {"query_class": "math", "model": "qwen2.5-coder:14b"}
|
||||
assert edit.references == ["trace-001", "trace-002"]
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
edit = Edit(**self._valid_edit_kwargs())
|
||||
as_json = edit.model_dump_json()
|
||||
restored = Edit.model_validate_json(as_json)
|
||||
|
||||
assert restored == edit
|
||||
|
||||
def test_pillar_must_be_valid_enum(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
kwargs = self._valid_edit_kwargs()
|
||||
kwargs["pillar"] = "not_a_pillar"
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
Edit(**kwargs)
|
||||
|
||||
def test_op_must_be_valid_enum(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
kwargs = self._valid_edit_kwargs()
|
||||
kwargs["op"] = "not_an_op"
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
Edit(**kwargs)
|
||||
|
||||
def test_payload_can_be_empty_dict(self) -> None:
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
kwargs = self._valid_edit_kwargs()
|
||||
kwargs["payload"] = {}
|
||||
|
||||
edit = Edit(**kwargs)
|
||||
assert edit.payload == {}
|
||||
|
||||
def test_references_default_empty_list(self) -> None:
|
||||
from openjarvis.learning.distillation.models import Edit
|
||||
|
||||
kwargs = self._valid_edit_kwargs()
|
||||
del kwargs["references"]
|
||||
|
||||
edit = Edit(**kwargs)
|
||||
assert edit.references == []
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FailureCluster
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFailureCluster:
|
||||
"""Tests for FailureCluster pydantic model."""
|
||||
|
||||
def _valid_cluster_kwargs(self) -> dict:
|
||||
return {
|
||||
"id": "cluster-001",
|
||||
"description": "Math word problems routed to qwen-3b",
|
||||
"sample_trace_ids": ["trace-001", "trace-002", "trace-003"],
|
||||
"student_failure_rate": 0.85,
|
||||
"teacher_success_rate": 0.95,
|
||||
"skill_gap": (
|
||||
"Student lacks chain-of-thought reasoning on multi-step arithmetic."
|
||||
),
|
||||
"addressed_by_edit_ids": ["edit-001", "edit-002"],
|
||||
}
|
||||
|
||||
def test_constructs_with_valid_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
cluster = FailureCluster(**self._valid_cluster_kwargs())
|
||||
|
||||
assert cluster.id == "cluster-001"
|
||||
assert cluster.student_failure_rate == 0.85
|
||||
assert cluster.teacher_success_rate == 0.95
|
||||
assert len(cluster.sample_trace_ids) == 3
|
||||
assert len(cluster.addressed_by_edit_ids) == 2
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
cluster = FailureCluster(**self._valid_cluster_kwargs())
|
||||
as_json = cluster.model_dump_json()
|
||||
restored = FailureCluster.model_validate_json(as_json)
|
||||
|
||||
assert restored == cluster
|
||||
|
||||
def test_addressed_by_edit_ids_defaults_empty(self) -> None:
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
kwargs = self._valid_cluster_kwargs()
|
||||
del kwargs["addressed_by_edit_ids"]
|
||||
|
||||
cluster = FailureCluster(**kwargs)
|
||||
assert cluster.addressed_by_edit_ids == []
|
||||
|
||||
def test_failure_rate_must_be_between_zero_and_one(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
kwargs = self._valid_cluster_kwargs()
|
||||
kwargs["student_failure_rate"] = 1.5
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
FailureCluster(**kwargs)
|
||||
|
||||
def test_success_rate_must_be_between_zero_and_one(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import FailureCluster
|
||||
|
||||
kwargs = self._valid_cluster_kwargs()
|
||||
kwargs["teacher_success_rate"] = -0.1
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
FailureCluster(**kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LearningPlan
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLearningPlan:
|
||||
"""Tests for LearningPlan pydantic model."""
|
||||
|
||||
def _valid_plan_kwargs(self) -> dict:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
FailureCluster,
|
||||
)
|
||||
|
||||
cluster = FailureCluster(
|
||||
id="cluster-001",
|
||||
description="Math routed to qwen-3b",
|
||||
sample_trace_ids=["t1", "t2", "t3"],
|
||||
student_failure_rate=0.8,
|
||||
teacher_success_rate=0.9,
|
||||
skill_gap="needs CoT",
|
||||
addressed_by_edit_ids=["edit-001"],
|
||||
)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
target="learning.routing.policy_map.math",
|
||||
payload={"query_class": "math", "model": "qwen2.5-coder:14b"},
|
||||
rationale="Math fails on small model",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
references=["t1"],
|
||||
)
|
||||
return {
|
||||
"session_id": "session-001",
|
||||
"diagnosis_summary": "## Diagnosis\nThe student misroutes math.",
|
||||
"failure_clusters": [cluster],
|
||||
"edits": [edit],
|
||||
"teacher_model": "claude-opus-4-6",
|
||||
"estimated_cost_usd": 1.42,
|
||||
"created_at": datetime(2026, 4, 8, 14, 22, 1, tzinfo=timezone.utc),
|
||||
}
|
||||
|
||||
def test_constructs_with_valid_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.models import LearningPlan
|
||||
|
||||
plan = LearningPlan(**self._valid_plan_kwargs())
|
||||
|
||||
assert plan.session_id == "session-001"
|
||||
assert len(plan.failure_clusters) == 1
|
||||
assert len(plan.edits) == 1
|
||||
assert plan.teacher_model == "claude-opus-4-6"
|
||||
assert plan.estimated_cost_usd == 1.42
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import LearningPlan
|
||||
|
||||
plan = LearningPlan(**self._valid_plan_kwargs())
|
||||
as_json = plan.model_dump_json()
|
||||
restored = LearningPlan.model_validate_json(as_json)
|
||||
|
||||
assert restored == plan
|
||||
|
||||
def test_empty_clusters_and_edits_allowed(self) -> None:
|
||||
# An aborted session may produce a plan with no clusters and no edits.
|
||||
from openjarvis.learning.distillation.models import LearningPlan
|
||||
|
||||
kwargs = self._valid_plan_kwargs()
|
||||
kwargs["failure_clusters"] = []
|
||||
kwargs["edits"] = []
|
||||
|
||||
plan = LearningPlan(**kwargs)
|
||||
assert plan.failure_clusters == []
|
||||
assert plan.edits == []
|
||||
|
||||
def test_estimated_cost_must_be_non_negative(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import LearningPlan
|
||||
|
||||
kwargs = self._valid_plan_kwargs()
|
||||
kwargs["estimated_cost_usd"] = -1.0
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
LearningPlan(**kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BenchmarkSnapshot
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBenchmarkSnapshot:
|
||||
"""Tests for BenchmarkSnapshot pydantic model."""
|
||||
|
||||
def _valid_snapshot_kwargs(self) -> dict:
|
||||
return {
|
||||
"benchmark_version": "personal_v3",
|
||||
"overall_score": 0.72,
|
||||
"cluster_scores": {"cluster-001": 0.65, "cluster-002": 0.80},
|
||||
"task_count": 50,
|
||||
"elapsed_seconds": 184.3,
|
||||
}
|
||||
|
||||
def test_constructs_with_valid_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
snap = BenchmarkSnapshot(**self._valid_snapshot_kwargs())
|
||||
|
||||
assert snap.benchmark_version == "personal_v3"
|
||||
assert snap.overall_score == 0.72
|
||||
assert snap.cluster_scores["cluster-001"] == 0.65
|
||||
assert snap.task_count == 50
|
||||
assert snap.elapsed_seconds == 184.3
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
snap = BenchmarkSnapshot(**self._valid_snapshot_kwargs())
|
||||
restored = BenchmarkSnapshot.model_validate_json(snap.model_dump_json())
|
||||
|
||||
assert restored == snap
|
||||
|
||||
def test_score_bounds(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
kwargs = self._valid_snapshot_kwargs()
|
||||
kwargs["overall_score"] = 1.5
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
BenchmarkSnapshot(**kwargs)
|
||||
|
||||
def test_task_count_must_be_non_negative(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
kwargs = self._valid_snapshot_kwargs()
|
||||
kwargs["task_count"] = -1
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
BenchmarkSnapshot(**kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# EditOutcome
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEditOutcome:
|
||||
"""Tests for EditOutcome pydantic model."""
|
||||
|
||||
def _valid_outcome_kwargs(self) -> dict:
|
||||
from datetime import datetime, timezone
|
||||
|
||||
return {
|
||||
"edit_id": "edit-001",
|
||||
"status": "applied",
|
||||
"benchmark_delta": 0.04,
|
||||
"cluster_deltas": {"cluster-001": 0.10, "cluster-002": 0.0},
|
||||
"error": None,
|
||||
"applied_at": datetime(2026, 4, 8, 14, 25, 0, tzinfo=timezone.utc),
|
||||
}
|
||||
|
||||
def test_applied_outcome(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOutcome
|
||||
|
||||
outcome = EditOutcome(**self._valid_outcome_kwargs())
|
||||
assert outcome.status == "applied"
|
||||
assert outcome.benchmark_delta == 0.04
|
||||
|
||||
def test_rejected_outcome_has_no_applied_at(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOutcome
|
||||
|
||||
outcome = EditOutcome(
|
||||
edit_id="edit-002",
|
||||
status="rejected_by_gate",
|
||||
benchmark_delta=-0.02,
|
||||
cluster_deltas={},
|
||||
error="regression: cluster-001 dropped 0.06",
|
||||
applied_at=None,
|
||||
)
|
||||
assert outcome.status == "rejected_by_gate"
|
||||
assert outcome.applied_at is None
|
||||
assert outcome.error is not None
|
||||
|
||||
def test_status_must_be_valid_literal(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import EditOutcome
|
||||
|
||||
kwargs = self._valid_outcome_kwargs()
|
||||
kwargs["status"] = "totally_made_up"
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
EditOutcome(**kwargs)
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOutcome
|
||||
|
||||
outcome = EditOutcome(**self._valid_outcome_kwargs())
|
||||
restored = EditOutcome.model_validate_json(outcome.model_dump_json())
|
||||
|
||||
assert restored == outcome
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# LearningSession
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestLearningSession:
|
||||
"""Tests for LearningSession pydantic model."""
|
||||
|
||||
def _valid_session_kwargs(self) -> dict:
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
BenchmarkSnapshot,
|
||||
SessionStatus,
|
||||
TriggerKind,
|
||||
)
|
||||
|
||||
snap = BenchmarkSnapshot(
|
||||
benchmark_version="personal_v1",
|
||||
overall_score=0.65,
|
||||
cluster_scores={"cluster-001": 0.50},
|
||||
task_count=30,
|
||||
elapsed_seconds=92.0,
|
||||
)
|
||||
return {
|
||||
"id": "session-001",
|
||||
"parent_session_id": None,
|
||||
"trigger": TriggerKind.SCHEDULED,
|
||||
"trigger_metadata": {"cron": "0 3 * * *"},
|
||||
"status": SessionStatus.INITIATED,
|
||||
"autonomy_mode": AutonomyMode.TIERED,
|
||||
"started_at": datetime(2026, 4, 8, 3, 0, 0, tzinfo=timezone.utc),
|
||||
"ended_at": None,
|
||||
"diagnosis_path": Path("/tmp/sessions/session-001/diagnosis.md"),
|
||||
"plan_path": Path("/tmp/sessions/session-001/plan.json"),
|
||||
"benchmark_before": snap,
|
||||
"benchmark_after": None,
|
||||
"edit_outcomes": [],
|
||||
"git_checkpoint_pre": "abc1234",
|
||||
"git_checkpoint_post": None,
|
||||
"teacher_cost_usd": 0.0,
|
||||
"error": None,
|
||||
}
|
||||
|
||||
def test_constructs_with_valid_fields(self) -> None:
|
||||
from openjarvis.learning.distillation.models import LearningSession
|
||||
|
||||
session = LearningSession(**self._valid_session_kwargs())
|
||||
assert session.id == "session-001"
|
||||
assert session.parent_session_id is None
|
||||
assert session.git_checkpoint_pre == "abc1234"
|
||||
assert session.benchmark_after is None
|
||||
|
||||
def test_round_trip_via_json(self) -> None:
|
||||
from openjarvis.learning.distillation.models import LearningSession
|
||||
|
||||
session = LearningSession(**self._valid_session_kwargs())
|
||||
as_json = session.model_dump_json()
|
||||
restored = LearningSession.model_validate_json(as_json)
|
||||
|
||||
assert restored == session
|
||||
|
||||
def test_supports_parent_session_chain(self) -> None:
|
||||
from openjarvis.learning.distillation.models import LearningSession
|
||||
|
||||
kwargs = self._valid_session_kwargs()
|
||||
kwargs["parent_session_id"] = "session-000"
|
||||
|
||||
session = LearningSession(**kwargs)
|
||||
assert session.parent_session_id == "session-000"
|
||||
|
||||
def test_status_must_be_valid_enum(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import LearningSession
|
||||
|
||||
kwargs = self._valid_session_kwargs()
|
||||
kwargs["status"] = "not_a_status"
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
LearningSession(**kwargs)
|
||||
|
||||
def test_teacher_cost_must_be_non_negative(self) -> None:
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from openjarvis.learning.distillation.models import LearningSession
|
||||
|
||||
kwargs = self._valid_session_kwargs()
|
||||
kwargs["teacher_cost_usd"] = -0.01
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
LearningSession(**kwargs)
|
||||
@@ -0,0 +1,174 @@
|
||||
"""Tests for DistillationOrchestrator — full session with mocks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
AutonomyMode,
|
||||
BenchmarkSnapshot,
|
||||
FailureCluster,
|
||||
SessionStatus,
|
||||
)
|
||||
from openjarvis.learning.distillation.triggers import OnDemandTrigger
|
||||
|
||||
|
||||
def _make_snapshot(overall: float = 0.6) -> BenchmarkSnapshot:
|
||||
return BenchmarkSnapshot(
|
||||
benchmark_version="personal_v1",
|
||||
overall_score=overall,
|
||||
cluster_scores={"c1": overall},
|
||||
task_count=10,
|
||||
elapsed_seconds=5.0,
|
||||
)
|
||||
|
||||
|
||||
def _make_diagnosis_result():
|
||||
from openjarvis.learning.distillation.diagnose.runner import DiagnosisResult
|
||||
|
||||
return DiagnosisResult(
|
||||
diagnosis_md="## Diagnosis\nMath routing is broken.",
|
||||
clusters=[
|
||||
FailureCluster(
|
||||
id="c1",
|
||||
description="Math routing",
|
||||
sample_trace_ids=["t1", "t2", "t3"],
|
||||
student_failure_rate=0.8,
|
||||
teacher_success_rate=0.95,
|
||||
skill_gap="needs CoT",
|
||||
)
|
||||
],
|
||||
cost_usd=0.05,
|
||||
tool_call_records=[],
|
||||
)
|
||||
|
||||
|
||||
def _make_mock_engine():
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = {
|
||||
"content": json.dumps(
|
||||
{
|
||||
"edits": [
|
||||
{
|
||||
"id": "edit-001",
|
||||
"pillar": "intelligence",
|
||||
"op": "set_model_for_query_class",
|
||||
"target": "routing.math",
|
||||
"payload": {
|
||||
"query_class": "math",
|
||||
"model": "qwen2.5-coder:14b",
|
||||
},
|
||||
"rationale": "Route math to bigger model",
|
||||
"expected_improvement": "c1",
|
||||
"risk_tier": "auto",
|
||||
"references": ["t1"],
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
"usage": {"total_tokens": 500},
|
||||
"cost_usd": 0.03,
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
return engine
|
||||
|
||||
|
||||
class TestDistillationOrchestrator:
|
||||
def test_full_session_completes(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.orchestrator import (
|
||||
DistillationOrchestrator,
|
||||
)
|
||||
|
||||
orch = DistillationOrchestrator(
|
||||
teacher_engine=_make_mock_engine(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(count=MagicMock(return_value=30)),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_store=MagicMock(),
|
||||
checkpoint_store=MagicMock(
|
||||
current_sha=MagicMock(return_value="abc123"),
|
||||
begin_stage=MagicMock(return_value=MagicMock(pre_stage_sha="abc123")),
|
||||
),
|
||||
openjarvis_home=tmp_path,
|
||||
autonomy_mode=AutonomyMode.AUTO,
|
||||
scorer=lambda **kw: _make_snapshot(0.65),
|
||||
benchmark_version="personal_v1",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"openjarvis.learning.distillation.orchestrator.DiagnosisRunner"
|
||||
) as MockDiag:
|
||||
MockDiag.return_value.run.return_value = _make_diagnosis_result()
|
||||
session = orch.run(OnDemandTrigger())
|
||||
|
||||
assert session.status in (
|
||||
SessionStatus.COMPLETED,
|
||||
SessionStatus.AWAITING_REVIEW,
|
||||
)
|
||||
assert session.teacher_cost_usd >= 0
|
||||
|
||||
def test_cold_start_returns_failed(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.orchestrator import (
|
||||
DistillationOrchestrator,
|
||||
)
|
||||
|
||||
trace_store = MagicMock()
|
||||
trace_store.count.return_value = 5 # Not enough
|
||||
|
||||
orch = DistillationOrchestrator(
|
||||
teacher_engine=MagicMock(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=trace_store,
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_store=MagicMock(),
|
||||
checkpoint_store=MagicMock(
|
||||
current_sha=MagicMock(return_value="abc123"),
|
||||
),
|
||||
openjarvis_home=tmp_path,
|
||||
autonomy_mode=AutonomyMode.TIERED,
|
||||
scorer=lambda **kw: _make_snapshot(),
|
||||
benchmark_version="personal_v1",
|
||||
)
|
||||
|
||||
session = orch.run(OnDemandTrigger())
|
||||
assert session.status == SessionStatus.FAILED
|
||||
assert "not enough traces" in (session.error or "").lower()
|
||||
|
||||
def test_session_persisted_to_store(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.orchestrator import (
|
||||
DistillationOrchestrator,
|
||||
)
|
||||
|
||||
session_store = MagicMock()
|
||||
|
||||
orch = DistillationOrchestrator(
|
||||
teacher_engine=_make_mock_engine(),
|
||||
teacher_model="claude-opus-4-6",
|
||||
trace_store=MagicMock(count=MagicMock(return_value=30)),
|
||||
benchmark_samples=[],
|
||||
student_runner=MagicMock(),
|
||||
judge=MagicMock(),
|
||||
session_store=session_store,
|
||||
checkpoint_store=MagicMock(
|
||||
current_sha=MagicMock(return_value="abc123"),
|
||||
begin_stage=MagicMock(return_value=MagicMock(pre_stage_sha="abc123")),
|
||||
),
|
||||
openjarvis_home=tmp_path,
|
||||
autonomy_mode=AutonomyMode.AUTO,
|
||||
scorer=lambda **kw: _make_snapshot(0.65),
|
||||
benchmark_version="personal_v1",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"openjarvis.learning.distillation.orchestrator.DiagnosisRunner"
|
||||
) as MockDiag:
|
||||
MockDiag.return_value.run.return_value = _make_diagnosis_result()
|
||||
orch.run(OnDemandTrigger())
|
||||
|
||||
assert session_store.save_session.called
|
||||
@@ -0,0 +1,107 @@
|
||||
"""Tests for openjarvis.learning.distillation.storage.paths module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_distillation_root
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestResolveDistillationRoot:
|
||||
"""Tests for resolve_distillation_root()."""
|
||||
|
||||
def test_default_is_under_home(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
monkeypatch.delenv("OPENJARVIS_HOME", raising=False)
|
||||
result = paths.resolve_distillation_root()
|
||||
assert result == Path.home() / ".openjarvis" / "learning"
|
||||
|
||||
def test_respects_openjarvis_home_env_var(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
custom = tmp_path / "custom_oj"
|
||||
monkeypatch.setenv("OPENJARVIS_HOME", str(custom))
|
||||
result = paths.resolve_distillation_root()
|
||||
assert result == custom / "learning"
|
||||
|
||||
def test_returns_absolute_path(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
monkeypatch.setenv("OPENJARVIS_HOME", str(tmp_path / "rel"))
|
||||
result = paths.resolve_distillation_root()
|
||||
assert result.is_absolute()
|
||||
|
||||
def test_rejects_path_inside_source_tree(
|
||||
self,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
# Find the OpenJarvis source root by walking up from the paths module.
|
||||
source_root = paths._find_source_root()
|
||||
assert source_root is not None # We must be running inside the repo.
|
||||
|
||||
# Force OPENJARVIS_HOME to point inside the source tree.
|
||||
monkeypatch.setenv("OPENJARVIS_HOME", str(source_root / "junk_dir"))
|
||||
|
||||
with pytest.raises(paths.ConfigurationError, match="inside the source tree"):
|
||||
paths.resolve_distillation_root()
|
||||
|
||||
def test_find_source_root_returns_repo_root(self) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
result = paths._find_source_root()
|
||||
assert result is not None
|
||||
assert (result / "pyproject.toml").exists()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ensure_distillation_dirs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEnsureDistillationDirs:
|
||||
"""Tests for ensure_distillation_dirs()."""
|
||||
|
||||
def test_creates_subdirs(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
monkeypatch.setenv("OPENJARVIS_HOME", str(tmp_path / "oj"))
|
||||
root = paths.ensure_distillation_dirs()
|
||||
|
||||
assert root.exists()
|
||||
assert (root / "sessions").exists()
|
||||
assert (root / "benchmarks").exists()
|
||||
assert (root / "benchmarks" / "reference_outputs").exists()
|
||||
assert (root / "pending_review").exists()
|
||||
|
||||
def test_idempotent(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from openjarvis.learning.distillation.storage import paths
|
||||
|
||||
monkeypatch.setenv("OPENJARVIS_HOME", str(tmp_path / "oj"))
|
||||
first = paths.ensure_distillation_dirs()
|
||||
second = paths.ensure_distillation_dirs()
|
||||
|
||||
assert first == second
|
||||
assert first.exists()
|
||||
@@ -0,0 +1,73 @@
|
||||
"""Tests for openjarvis.learning.distillation.pending_queue module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
|
||||
|
||||
def _make_edit(edit_id: str = "edit-001") -> Edit:
|
||||
return Edit(
|
||||
id=edit_id,
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.REPLACE_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"new_content": "New prompt.\n"},
|
||||
rationale="Better prompt",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
|
||||
|
||||
class TestPendingQueue:
|
||||
def test_enqueue_creates_file(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
|
||||
queue = PendingQueue(tmp_path / "pending_review")
|
||||
queue.enqueue("session-001", _make_edit())
|
||||
files = list((tmp_path / "pending_review").glob("*.json"))
|
||||
assert len(files) == 1
|
||||
assert "session-001" in files[0].name
|
||||
assert "edit-001" in files[0].name
|
||||
|
||||
def test_list_pending(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
|
||||
queue = PendingQueue(tmp_path / "pending_review")
|
||||
queue.enqueue("session-001", _make_edit("e1"))
|
||||
queue.enqueue("session-001", _make_edit("e2"))
|
||||
pending = queue.list_pending()
|
||||
assert len(pending) == 2
|
||||
ids = {p["edit"]["id"] for p in pending}
|
||||
assert ids == {"e1", "e2"}
|
||||
|
||||
def test_resolve_removes_file(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
|
||||
queue = PendingQueue(tmp_path / "pending_review")
|
||||
queue.enqueue("session-001", _make_edit())
|
||||
assert len(queue.list_pending()) == 1
|
||||
queue.resolve("session-001", "edit-001")
|
||||
assert len(queue.list_pending()) == 0
|
||||
|
||||
def test_list_empty_queue(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
|
||||
queue = PendingQueue(tmp_path / "pending_review")
|
||||
assert queue.list_pending() == []
|
||||
|
||||
def test_get_pending_edit(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.pending_queue import PendingQueue
|
||||
|
||||
queue = PendingQueue(tmp_path / "pending_review")
|
||||
queue.enqueue("session-001", _make_edit())
|
||||
edit_data = queue.get("session-001", "edit-001")
|
||||
assert edit_data is not None
|
||||
assert edit_data["edit"]["id"] == "edit-001"
|
||||
assert edit_data["session_id"] == "session-001"
|
||||
@@ -0,0 +1,272 @@
|
||||
"""Tests for openjarvis.learning.distillation.plan.planner module.
|
||||
|
||||
All tests use mocked CloudEngine — no live API calls.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from openjarvis.learning.distillation.models import (
|
||||
EditOp,
|
||||
EditRiskTier,
|
||||
FailureCluster,
|
||||
)
|
||||
|
||||
|
||||
def _make_clusters() -> list[FailureCluster]:
|
||||
return [
|
||||
FailureCluster(
|
||||
id="cluster-001",
|
||||
description="Math queries routed to qwen-3b",
|
||||
sample_trace_ids=["t1", "t2", "t3"],
|
||||
student_failure_rate=0.8,
|
||||
teacher_success_rate=0.95,
|
||||
skill_gap="Student lacks chain-of-thought on multi-step math",
|
||||
),
|
||||
FailureCluster(
|
||||
id="cluster-002",
|
||||
description="Calculator tool not used",
|
||||
sample_trace_ids=["t4", "t5", "t6"],
|
||||
student_failure_rate=0.6,
|
||||
teacher_success_rate=0.9,
|
||||
skill_gap="Student does not invoke calculator",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _make_teacher_response(edits_json: list[dict]) -> dict:
|
||||
"""Create a mock engine.generate() response with edit list."""
|
||||
return {
|
||||
"content": json.dumps({"edits": edits_json}),
|
||||
"usage": {"prompt_tokens": 500, "completion_tokens": 300, "total_tokens": 800},
|
||||
"cost_usd": 0.03,
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
|
||||
|
||||
def _make_edit_dict(
|
||||
edit_id: str = "edit-001",
|
||||
pillar: str = "intelligence",
|
||||
op: str = "set_model_for_query_class",
|
||||
target: str = "learning.routing.policy_map.math",
|
||||
payload: dict | None = None,
|
||||
expected_improvement: str = "cluster-001",
|
||||
) -> dict:
|
||||
return {
|
||||
"id": edit_id,
|
||||
"pillar": pillar,
|
||||
"op": op,
|
||||
"target": target,
|
||||
"payload": payload or {"query_class": "math", "model": "qwen2.5-coder:14b"},
|
||||
"rationale": "Route math queries to a bigger model",
|
||||
"expected_improvement": expected_improvement,
|
||||
"risk_tier": "auto",
|
||||
"references": ["t1", "t2"],
|
||||
}
|
||||
|
||||
|
||||
class TestLearningPlanner:
|
||||
"""Tests for LearningPlanner."""
|
||||
|
||||
def test_produces_learning_plan(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_teacher_response([_make_edit_dict()])
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=tmp_path / "session-001",
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md="## Diagnosis\nMath routing is broken.",
|
||||
clusters=_make_clusters(),
|
||||
)
|
||||
|
||||
assert plan.session_id == "session-001"
|
||||
assert len(plan.edits) == 1
|
||||
assert plan.edits[0].op == EditOp.SET_MODEL_FOR_QUERY_CLASS
|
||||
assert plan.teacher_model == "claude-opus-4-6"
|
||||
|
||||
def test_assigns_risk_tiers(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
# Teacher incorrectly sets MANUAL for an auto-tier op
|
||||
edit_dict = _make_edit_dict()
|
||||
edit_dict["risk_tier"] = "manual"
|
||||
engine.generate.return_value = _make_teacher_response([edit_dict])
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=tmp_path / "session-001",
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=_make_clusters(),
|
||||
)
|
||||
|
||||
# Should be overwritten to AUTO
|
||||
assert plan.edits[0].risk_tier == EditRiskTier.AUTO
|
||||
|
||||
def test_persists_plan_json(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_teacher_response([_make_edit_dict()])
|
||||
|
||||
session_dir = tmp_path / "session-001"
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=session_dir,
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=_make_clusters(),
|
||||
)
|
||||
|
||||
plan_path = session_dir / "plan.json"
|
||||
assert plan_path.exists()
|
||||
data = json.loads(plan_path.read_text())
|
||||
assert data["session_id"] == "session-001"
|
||||
|
||||
def test_drops_cluster_with_zero_rates(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_teacher_response([_make_edit_dict()])
|
||||
|
||||
clusters = [
|
||||
FailureCluster(
|
||||
id="cluster-bad",
|
||||
description="No evidence",
|
||||
sample_trace_ids=[],
|
||||
student_failure_rate=0.0,
|
||||
teacher_success_rate=0.0,
|
||||
skill_gap="Speculative",
|
||||
),
|
||||
FailureCluster(
|
||||
id="cluster-good",
|
||||
description="Real evidence",
|
||||
sample_trace_ids=["t1", "t2", "t3"],
|
||||
student_failure_rate=0.8,
|
||||
teacher_success_rate=0.9,
|
||||
skill_gap="Verified gap",
|
||||
),
|
||||
]
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=tmp_path / "session-001",
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=clusters,
|
||||
)
|
||||
|
||||
# cluster-bad should be dropped (marked in skill_gap)
|
||||
bad = next(c for c in plan.failure_clusters if c.id == "cluster-bad")
|
||||
assert "dropped" in bad.skill_gap.lower()
|
||||
assert bad.addressed_by_edit_ids == []
|
||||
# cluster-good should survive
|
||||
good = next(c for c in plan.failure_clusters if c.id == "cluster-good")
|
||||
assert "dropped" not in good.skill_gap.lower()
|
||||
|
||||
def test_all_clusters_dropped_returns_empty_edits(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_teacher_response([])
|
||||
|
||||
clusters = [
|
||||
FailureCluster(
|
||||
id="cluster-bad",
|
||||
description="No evidence",
|
||||
sample_trace_ids=[],
|
||||
student_failure_rate=0.0,
|
||||
teacher_success_rate=0.0,
|
||||
skill_gap="Speculative",
|
||||
),
|
||||
]
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=tmp_path / "session-001",
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=clusters,
|
||||
)
|
||||
|
||||
assert plan.edits == []
|
||||
assert plan.failure_clusters[0].addressed_by_edit_ids == []
|
||||
|
||||
def test_persists_teacher_trace_jsonl(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = _make_teacher_response([_make_edit_dict()])
|
||||
|
||||
session_dir = tmp_path / "session-001"
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=session_dir,
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=_make_clusters(),
|
||||
)
|
||||
|
||||
jsonl_path = session_dir / "teacher_traces" / "plan.jsonl"
|
||||
assert jsonl_path.exists()
|
||||
lines = jsonl_path.read_text().strip().splitlines()
|
||||
assert len(lines) >= 1
|
||||
record = json.loads(lines[0])
|
||||
assert "cost_usd" in record
|
||||
|
||||
def test_handles_malformed_teacher_output(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.plan.planner import LearningPlanner
|
||||
|
||||
engine = MagicMock()
|
||||
engine.generate.return_value = {
|
||||
"content": "This is not valid JSON at all",
|
||||
"usage": {"total_tokens": 100},
|
||||
"cost_usd": 0.01,
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
|
||||
planner = LearningPlanner(
|
||||
teacher_engine=engine,
|
||||
teacher_model="claude-opus-4-6",
|
||||
session_id="session-001",
|
||||
session_dir=tmp_path / "session-001",
|
||||
prompt_reader=lambda t: "",
|
||||
)
|
||||
plan = planner.run(
|
||||
diagnosis_md="## Diagnosis",
|
||||
clusters=_make_clusters(),
|
||||
)
|
||||
|
||||
# Should return a plan with no edits rather than crashing
|
||||
assert plan.edits == []
|
||||
@@ -0,0 +1,212 @@
|
||||
"""Tests for openjarvis.learning.distillation.plan.prompt_diff module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class TestChangedLineRatio:
|
||||
"""Tests for changed_line_ratio()."""
|
||||
|
||||
def test_identical_strings(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
changed_line_ratio,
|
||||
)
|
||||
|
||||
assert changed_line_ratio("hello\nworld\n", "hello\nworld\n") == 0.0
|
||||
|
||||
def test_completely_different(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
changed_line_ratio,
|
||||
)
|
||||
|
||||
ratio = changed_line_ratio("aaa\nbbb\nccc\n", "xxx\nyyy\nzzz\n")
|
||||
assert ratio == 1.0
|
||||
|
||||
def test_partial_change(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
changed_line_ratio,
|
||||
)
|
||||
|
||||
original = "line1\nline2\nline3\nline4\n"
|
||||
modified = "line1\nchanged\nline3\nline4\n"
|
||||
ratio = changed_line_ratio(original, modified)
|
||||
# 1 out of 4 lines changed
|
||||
assert 0.2 <= ratio <= 0.35
|
||||
|
||||
def test_empty_original(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
changed_line_ratio,
|
||||
)
|
||||
|
||||
# Adding to empty = 100% change
|
||||
ratio = changed_line_ratio("", "new content\n")
|
||||
assert ratio == 1.0
|
||||
|
||||
def test_empty_both(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
changed_line_ratio,
|
||||
)
|
||||
|
||||
assert changed_line_ratio("", "") == 0.0
|
||||
|
||||
|
||||
class TestApplyUnifiedDiff:
|
||||
"""Tests for apply_unified_diff()."""
|
||||
|
||||
def test_applies_simple_patch(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
apply_unified_diff,
|
||||
)
|
||||
|
||||
original = "line1\nline2\nline3\n"
|
||||
diff = (
|
||||
"--- a/prompt.md\n"
|
||||
"+++ b/prompt.md\n"
|
||||
"@@ -1,3 +1,3 @@\n"
|
||||
" line1\n"
|
||||
"-line2\n"
|
||||
"+changed_line2\n"
|
||||
" line3\n"
|
||||
)
|
||||
result = apply_unified_diff(original, diff)
|
||||
assert result == "line1\nchanged_line2\nline3\n"
|
||||
|
||||
def test_returns_none_on_bad_diff(self) -> None:
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
apply_unified_diff,
|
||||
)
|
||||
|
||||
result = apply_unified_diff("hello\n", "not a valid diff")
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestMaybeDowngradeToReplace:
|
||||
"""Tests for maybe_downgrade_to_replace()."""
|
||||
|
||||
def test_non_patch_op_passes_through(self) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
maybe_downgrade_to_replace,
|
||||
)
|
||||
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
target="routing.math",
|
||||
payload={"model": "qwen2.5-coder:14b"},
|
||||
rationale="test",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO,
|
||||
)
|
||||
result = maybe_downgrade_to_replace(edit, prompt_reader=lambda t: "")
|
||||
assert result.op == EditOp.SET_MODEL_FOR_QUERY_CLASS
|
||||
|
||||
def test_small_diff_stays_patch(self, tmp_path: Path) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
maybe_downgrade_to_replace,
|
||||
)
|
||||
|
||||
original = (
|
||||
"line1\nline2\nline3\nline4\nline5\nline6\nline7\nline8\nline9\nline10\n"
|
||||
)
|
||||
diff = (
|
||||
"--- a/prompt.md\n"
|
||||
"+++ b/prompt.md\n"
|
||||
"@@ -2,1 +2,1 @@\n"
|
||||
" line1\n"
|
||||
"-line2\n"
|
||||
"+changed\n"
|
||||
" line3\n"
|
||||
)
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": diff},
|
||||
rationale="Small fix",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
result = maybe_downgrade_to_replace(edit, prompt_reader=lambda t: original)
|
||||
assert result.op == EditOp.PATCH_SYSTEM_PROMPT
|
||||
|
||||
def test_large_diff_downgrades_to_replace(self) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
maybe_downgrade_to_replace,
|
||||
)
|
||||
|
||||
original = "old1\nold2\nold3\nold4\n"
|
||||
diff = (
|
||||
"--- a/prompt.md\n"
|
||||
"+++ b/prompt.md\n"
|
||||
"@@ -1,4 +1,4 @@\n"
|
||||
"-old1\n"
|
||||
"-old2\n"
|
||||
"-old3\n"
|
||||
"-old4\n"
|
||||
"+new1\n"
|
||||
"+new2\n"
|
||||
"+new3\n"
|
||||
"+new4\n"
|
||||
)
|
||||
edit = Edit(
|
||||
id="edit-002",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": diff},
|
||||
rationale="Major rewrite",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
result = maybe_downgrade_to_replace(edit, prompt_reader=lambda t: original)
|
||||
assert result.op == EditOp.REPLACE_SYSTEM_PROMPT
|
||||
assert "new_content" in result.payload
|
||||
assert "new1" in result.payload["new_content"]
|
||||
|
||||
def test_bad_diff_downgrades_to_replace_with_raw(self) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.prompt_diff import (
|
||||
maybe_downgrade_to_replace,
|
||||
)
|
||||
|
||||
edit = Edit(
|
||||
id="edit-003",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": "this is not a valid unified diff"},
|
||||
rationale="Bad diff",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.REVIEW,
|
||||
)
|
||||
result = maybe_downgrade_to_replace(
|
||||
edit, prompt_reader=lambda t: "original content\n"
|
||||
)
|
||||
# Can't apply the diff, so it should downgrade
|
||||
assert result.op == EditOp.REPLACE_SYSTEM_PROMPT
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Tests for openjarvis.learning.distillation.gate.regression module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from openjarvis.learning.distillation.models import BenchmarkSnapshot
|
||||
|
||||
|
||||
def _make_snapshot(
|
||||
overall: float = 0.7,
|
||||
clusters: dict[str, float] | None = None,
|
||||
) -> BenchmarkSnapshot:
|
||||
return BenchmarkSnapshot(
|
||||
benchmark_version="personal_v1",
|
||||
overall_score=overall,
|
||||
cluster_scores=clusters or {"c1": 0.6, "c2": 0.8},
|
||||
task_count=50,
|
||||
elapsed_seconds=60.0,
|
||||
)
|
||||
|
||||
|
||||
class TestRegressionCheck:
|
||||
"""Tests for regression_check()."""
|
||||
|
||||
def test_no_regression_when_all_improve(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.6, clusters={"c1": 0.5, "c2": 0.7})
|
||||
after = _make_snapshot(overall=0.7, clusters={"c1": 0.6, "c2": 0.8})
|
||||
result = regression_check(before, after, max_regression=0.05)
|
||||
assert not result.has_regression
|
||||
|
||||
def test_detects_cluster_regression(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.7, clusters={"c1": 0.6, "c2": 0.8})
|
||||
after = _make_snapshot(overall=0.72, clusters={"c1": 0.65, "c2": 0.70})
|
||||
result = regression_check(before, after, max_regression=0.05)
|
||||
assert result.has_regression
|
||||
assert "c2" in result.regressed_clusters
|
||||
|
||||
def test_small_drop_within_threshold(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.7, clusters={"c1": 0.6, "c2": 0.8})
|
||||
after = _make_snapshot(overall=0.72, clusters={"c1": 0.63, "c2": 0.76})
|
||||
result = regression_check(before, after, max_regression=0.05)
|
||||
assert not result.has_regression
|
||||
|
||||
def test_new_cluster_in_after_not_flagged(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.7, clusters={"c1": 0.6})
|
||||
after = _make_snapshot(overall=0.75, clusters={"c1": 0.65, "c2": 0.8})
|
||||
result = regression_check(before, after, max_regression=0.05)
|
||||
assert not result.has_regression
|
||||
|
||||
def test_missing_cluster_in_after_flagged(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.7, clusters={"c1": 0.6, "c2": 0.8})
|
||||
after = _make_snapshot(overall=0.72, clusters={"c1": 0.65})
|
||||
# c2 disappeared — treat as regression (score went from 0.8 to 0.0)
|
||||
result = regression_check(before, after, max_regression=0.05)
|
||||
assert result.has_regression
|
||||
assert "c2" in result.regressed_clusters
|
||||
|
||||
def test_result_has_details(self) -> None:
|
||||
from openjarvis.learning.distillation.gate.regression import (
|
||||
regression_check,
|
||||
)
|
||||
|
||||
before = _make_snapshot(overall=0.7, clusters={"c1": 0.6, "c2": 0.8})
|
||||
after = _make_snapshot(overall=0.68, clusters={"c1": 0.55, "c2": 0.75})
|
||||
result = regression_check(before, after, max_regression=0.03)
|
||||
assert result.has_regression
|
||||
assert len(result.regressed_clusters) >= 1
|
||||
# Check that deltas are provided
|
||||
for cluster_id, delta in result.regressed_clusters.items():
|
||||
assert delta < 0
|
||||
@@ -0,0 +1,120 @@
|
||||
"""Tests for openjarvis.learning.distillation.plan.risk_tier module."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class TestTierTable:
|
||||
"""Tests for TIER_TABLE completeness."""
|
||||
|
||||
def test_every_edit_op_has_a_tier(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
from openjarvis.learning.distillation.plan.risk_tier import TIER_TABLE
|
||||
|
||||
for op in EditOp:
|
||||
assert op in TIER_TABLE, f"Missing tier for {op}"
|
||||
|
||||
def test_no_extra_keys(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp
|
||||
from openjarvis.learning.distillation.plan.risk_tier import TIER_TABLE
|
||||
|
||||
for key in TIER_TABLE:
|
||||
assert key in EditOp, f"Extra key in TIER_TABLE: {key}"
|
||||
|
||||
|
||||
class TestAssignTier:
|
||||
"""Tests for assign_tier() function."""
|
||||
|
||||
def test_intelligence_ops_are_auto(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.SET_MODEL_FOR_QUERY_CLASS) == EditRiskTier.AUTO
|
||||
assert assign_tier(EditOp.SET_MODEL_PARAM) == EditRiskTier.AUTO
|
||||
|
||||
def test_tool_ops_are_auto(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.ADD_TOOL_TO_AGENT) == EditRiskTier.AUTO
|
||||
assert assign_tier(EditOp.REMOVE_TOOL_FROM_AGENT) == EditRiskTier.AUTO
|
||||
assert assign_tier(EditOp.EDIT_TOOL_DESCRIPTION) == EditRiskTier.AUTO
|
||||
|
||||
def test_agent_param_is_auto(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.SET_AGENT_PARAM) == EditRiskTier.AUTO
|
||||
|
||||
def test_prompt_ops_are_review(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.PATCH_SYSTEM_PROMPT) == EditRiskTier.REVIEW
|
||||
assert assign_tier(EditOp.REPLACE_SYSTEM_PROMPT) == EditRiskTier.REVIEW
|
||||
|
||||
def test_agent_class_is_review(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.SET_AGENT_CLASS) == EditRiskTier.REVIEW
|
||||
|
||||
def test_few_shot_is_review(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.EDIT_FEW_SHOT_EXEMPLARS) == EditRiskTier.REVIEW
|
||||
|
||||
def test_lora_is_manual(self) -> None:
|
||||
from openjarvis.learning.distillation.models import EditOp, EditRiskTier
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tier
|
||||
|
||||
assert assign_tier(EditOp.LORA_FINETUNE) == EditRiskTier.MANUAL
|
||||
|
||||
|
||||
class TestAssignTiers:
|
||||
"""Tests for assign_tiers() batch function."""
|
||||
|
||||
def test_overwrites_teacher_tier(self) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tiers
|
||||
|
||||
# Teacher incorrectly sets AUTO for a prompt edit
|
||||
edit = Edit(
|
||||
id="edit-001",
|
||||
pillar=EditPillar.AGENT,
|
||||
op=EditOp.PATCH_SYSTEM_PROMPT,
|
||||
target="agents.simple.system_prompt",
|
||||
payload={"diff": "--- a\n+++ b\n@@ -1 +1 @@\n-old\n+new"},
|
||||
rationale="Improve prompt",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO, # Wrong — should be REVIEW
|
||||
)
|
||||
result = assign_tiers([edit])
|
||||
assert result[0].risk_tier == EditRiskTier.REVIEW
|
||||
|
||||
def test_preserves_correct_tier(self) -> None:
|
||||
from openjarvis.learning.distillation.models import (
|
||||
Edit,
|
||||
EditOp,
|
||||
EditPillar,
|
||||
EditRiskTier,
|
||||
)
|
||||
from openjarvis.learning.distillation.plan.risk_tier import assign_tiers
|
||||
|
||||
edit = Edit(
|
||||
id="edit-002",
|
||||
pillar=EditPillar.INTELLIGENCE,
|
||||
op=EditOp.SET_MODEL_FOR_QUERY_CLASS,
|
||||
target="learning.routing.policy_map.math",
|
||||
payload={"query_class": "math", "model": "qwen2.5-coder:14b"},
|
||||
rationale="Route math to bigger model",
|
||||
expected_improvement="cluster-001",
|
||||
risk_tier=EditRiskTier.AUTO, # Correct
|
||||
)
|
||||
result = assign_tiers([edit])
|
||||
assert result[0].risk_tier == EditRiskTier.AUTO
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user