#include "strata/core/conversation_snapshot.hpp" #include "strata/kernels/kv_q4.hpp" #include #include #include #include #include #include using namespace strata::core; using namespace strata::kernels; namespace { int checks = 0; void check(bool ok, const char* label) { ++checks; if (!ok) { std::fprintf(stderr, "FAIL: %s\n", label); std::exit(1); } } void cuda_check(cudaError_t e) { if (e != cudaSuccess) { std::fprintf(stderr, "CUDA: %s\n", cudaGetErrorString(e)); std::exit(1); } } struct Fixture { ModelGeometry g; QsaState state; std::vector device, host; std::array sources{}; std::array sizes{}; template void alloc(T*& p, size_t n, bool pinned = false) { if (!n) return; void* raw = nullptr; if (pinned) { cuda_check(cudaHostAlloc(&raw, n, cudaHostAllocMapped)); host.push_back(raw); void* mapped = nullptr; cuda_check(cudaHostGetDevicePointer(&mapped, raw, 0)); p = static_cast(mapped); } else { cuda_check(cudaMalloc(&raw, n)); device.push_back(raw); p = static_cast(raw); } } Fixture(int fmt, int mode) { auto& st = state; st.kv_mode = mode; st.kv_int8 = fmt==kKvInt8; st.kv_q4 = fmt==kKvQ4; st.kv_hybrid = fmt==3; st.n_pages = 24; st.n_slots = mode ? 4 : st.n_pages; st.max_cells = st.n_pages * 4; st.idx_pooled_rows = mode==2 ? 2 : st.max_cells/4+2; const size_t per = fmt==kKvQ4 ? kv_q4_bytes_per_head((int)g.head_dim) : g.head_dim*((st.kv_int8 || st.kv_hybrid) ? 1:2); const size_t rows = st.max_cells*g.n_head_kv, slot_rows = st.n_slots*4*g.n_head_kv; sizes = {rows*per, rows*per, fmt==kKvInt8 ? rows*(g.head_dim/64)*2:0, fmt==kKvInt8 ? rows*(g.head_dim/64)*2:0, (size_t)st.idx_pooled_rows*g.idx_key_dim*4}; if (fmt==3) { sizes = {rows*per, rows*kv_q4_bytes_per_head((int)g.head_dim), rows*(g.head_dim/64)*2, 0, (size_t)st.idx_pooled_rows*g.idx_key_dim*4}; alloc(st.k_q,sizes[0]); alloc(st.v_q4,sizes[1]); alloc(st.k_scale,sizes[2]); sources[0]=st.k_q; sources[1]=st.v_q4; sources[2]=st.k_scale; } else if (fmt==kKvQ4) { alloc(st.k_q4,slot_rows*per); alloc(st.v_q4,slot_rows*per); if (mode) {alloc(st.host.k_q4,sizes[0],true); alloc(st.host.v_q4,sizes[1],true);} sources[0]=mode?st.host.k_q4:st.k_q4; sources[1]=mode?st.host.v_q4:st.v_q4; } else if (fmt==kKvInt8) { alloc(st.k_q,slot_rows*per); alloc(st.v_q,slot_rows*per); alloc(st.k_scale,slot_rows*(g.head_dim/64)*2); alloc(st.v_scale,slot_rows*(g.head_dim/64)*2); if (mode) { alloc(st.host.k_q,sizes[0],true); alloc(st.host.v_q,sizes[1],true); alloc(st.host.k_scale,sizes[2],true); alloc(st.host.v_scale,sizes[3],true); } sources[0]=mode?st.host.k_q:st.k_q; sources[1]=mode?st.host.v_q:st.v_q; sources[2]=mode?st.host.k_scale:st.k_scale; sources[3]=mode?st.host.v_scale:st.v_scale; } else { alloc(st.k_pool,slot_rows*per); alloc(st.v_pool,slot_rows*per); if (mode) {alloc(st.host.k_pool,sizes[0],true); alloc(st.host.v_pool,sizes[1],true);} sources[0]=mode?st.host.k_pool:st.k_pool; sources[1]=mode?st.host.v_pool:st.v_pool; } alloc(st.idx_pooled,sizes[4]); sources[4]=st.idx_pooled; alloc(st.page_table,st.n_pages*4); if (mode==1) { auto& m=st.map; m.page_table=st.page_table; m.n_blocks=st.n_pages; m.n_slots=st.n_slots; alloc(m.slot_block,st.n_slots*4); alloc(m.slot_stamp,st.n_slots*4); alloc(m.slot_ref,st.n_slots*4); alloc(m.miss_block,st.n_slots*4); alloc(m.miss_slot,st.n_slots*4); alloc(m.ctl,kKvCtlInts*4); } else if (mode==2) { kv_ring_table(st.page_table,st.n_pages,st.n_slots,nullptr); } } void fill(uint8_t salt) { for (size_t i=0;i data(sizes[i]); for (size_t j=0;j offset) cuda_check(cudaMemset(static_cast(sources[i])+offset, salt, sizes[i]-offset)); } } ~Fixture() { for (void* p:device) cudaFree(p); for (void* p:host) cudaFreeHost(p); } }; bool equal(const ConversationKv& a,const ConversationKv& b) { return a.k==b.k && a.v==b.v && a.k_scale==b.k_scale && a.v_scale==b.v_scale && a.pooled==b.pooled; } void full_session(int fmt, int mode, int experts) { Fixture main(fmt,mode), draft(fmt==3?kKvInt8:fmt,2); auto& g=main.g; g.n_layers=4; g.n_expert=experts; g.ssm_state_size=2; g.ssm_v_heads=2; g.ssm_conv_channels=8; SessionState ss; ss.max_cells=96; ss.qsa_states=&main.state; ss.layer_hi=g.n_layers; ss.gdn_alloc=g.n_gdn_layers(); ss.qsa_alloc=g.n_qsa_layers(); // the whole-model carve std::string err; ConversationStateSizes sizes; check(conversation_state_sizes(g,sizes,err),"whole-session geometry sizes"); main.alloc(ss.gdn_state,sizes.gdn); main.alloc(ss.ple_hist,sizes.ple); main.alloc(main.state.idx_tail,sizes.tail); main.alloc(main.state.idx_dead,sizes.dead); main.alloc(main.state.idx_block_pos,sizes.block_pos); std::vector ids(65); for (size_t i=0;i images; std::vector checkpoints; auto fill=[&](uint8_t salt) { main.fill(salt); draft.fill(salt); for (const auto& [p,n] : std::vector>{ {ss.gdn_state,sizes.gdn},{ss.ple_hist,sizes.ple},{main.state.idx_tail,sizes.tail}, {main.state.idx_dead,sizes.dead},{main.state.idx_block_pos,sizes.block_pos}}) cuda_check(cudaMemset(p,salt,n)); // A real indexer maintains its spare pooled row as a copy of idx_dead. cuda_check(cudaMemcpy(main.state.idx_pooled+(ids.size()/4)*g.idx_key_dim, main.state.idx_dead,sizes.dead,cudaMemcpyDeviceToDevice)); cuda_check(cudaDeviceSynchronize()); }; fill(13); ConversationCheckpoint checkpoint; checkpoint.ids.assign(ids.begin(),ids.begin()+3); check(conversation_checkpoint_save(checkpoint,ss,g,err),"save complete running checkpoint"); checkpoint.used = 17; checkpoints.push_back(std::move(checkpoint)); const ConversationView view{ids,images,checkpoints,true}; SavedConversation a,b,restored; check(conversation_snapshot_save(a,view,ss,g,draft.state,err),"capture complete A"); check(a.checkpoints[0].used == 17,"upstream checkpoint LRU stamp survives capture"); fill(177); check(conversation_snapshot_save(b,view,ss,g,draft.state,err),"capture complete B"); check(a.live.dead!=b.live.dead,"different opening-state spare keys in regression fixture"); auto bad=a; bad.kv.back().k.pop_back(); check(conversation_snapshot_restore(bad,ss,g,draft.state,err)==ConversationRestore::invalid, "reject invalid late draft before any main-layer write"); check(conversation_snapshot_save(restored,view,ss,g,draft.state,err),"capture B after refused restore"); check(restored.live.gdn==b.live.gdn && restored.live.dead==b.live.dead && equal(restored.kv[0],b.kv[0]) && equal(restored.kv[1],b.kv[1]),"refusal preserves all tested state"); check(conversation_snapshot_restore(a,ss,g,draft.state,err)==ConversationRestore::restored,"restore complete A"); check(conversation_snapshot_save(restored,view,ss,g,draft.state,err),"capture restored A"); check(restored.live.gdn==a.live.gdn && restored.live.ple==a.live.ple && restored.live.tails==a.live.tails && restored.live.dead==a.live.dead && restored.live.block_pos==a.live.block_pos && equal(restored.kv[0],a.kv[0]) && equal(restored.kv[1],a.kv[1]),"whole-session A/B/A exactness including spare key"); check(ss.ple_prev[0]==64 && ss.ple_prev[1]==65,"PLE token window reconstructed"); for (int64_t dirty : {65, 3, 0}) { check(conversation_snapshot_restore(a,ss,g,draft.state,err)==ConversationRestore::restored,"restore growth fixture base"); ConversationKvReuse reuse{a.kv,65,dirty}; const uint8_t* original = nullptr; reuse.kv[0].k.visit(0,1,[&](const uint8_t* p,size_t,size_t){original=p;return true;}); main.fill_after(91,dirty); draft.fill_after(91,std::max(0,dirty-1)); cuda_check(cudaMemset(ss.gdn_state,91,sizes.gdn)); ids.resize(70); for (size_t i=65;i=4)==(reused>0),"only complete unchanged pages or rows are retained"); incremental.kv[0].k.visit(0,1,[&](const uint8_t* p,size_t,size_t){ check(p==original,"growth never reallocates the retained payload");return true; }); check(conversation_snapshot_restore(incremental,ss,g,draft.state,err)==ConversationRestore::restored,"restore segmented incremental snapshot"); uint64_t fingerprint=0; check(conversation_kv_verify(incremental.kv.back(),draft.state,g,70,false,fingerprint,err),"incremental draft authoritative and ring read-back"); ids.resize(65); } check(conversation_checkpoint_restore(a.checkpoints[0],ss,g,err),"restore early running checkpoint"); std::vector spare(sizes.dead); cuda_check(cudaMemcpy(spare.data(),main.state.idx_pooled,sizes.dead,cudaMemcpyDeviceToHost)); check(spare==a.checkpoints[0].dead,"checkpoint rebuilds spare row over a later completed block"); check(ss.ple_prev[0]==2 && ss.ple_prev[1]==3,"checkpoint PLE token window"); } } int main() { int devices=0; if (cudaGetDeviceCount(&devices)!=cudaSuccess || !devices) return 77; for (int fmt : std::array{kKvF16,kKvInt8,kKvQ4,3}) for (int mode : {0,1,2}) { if (fmt==3 && mode!=0) continue; Fixture f(fmt,mode); if (fmt==3) { std::string err; ConversationKv invalid; for (int unsupported : {1,2}) { f.state.kv_mode=unsupported; check(conversation_kv_bytes(f.state,f.g,9,true)==0,"hybrid streaming/ring rejected without block-mover abort"); check(!conversation_kv_save(invalid,f.state,f.g,9,true,err),"hybrid streaming/ring capture fails closed"); } f.state.kv_mode=0; } { std::string err; ConversationKv invalid; for (const int64_t bad : {-1LL, 97LL, (long long) std::numeric_limits::max()}) { check(!conversation_kv_save(invalid,f.state,f.g,bad,mode!=2,err),"reject invalid save extent before allocation"); check(!conversation_kv_restore(invalid,f.state,f.g,bad,mode!=2,err),"reject invalid restore extent before copying"); check(conversation_kv_bytes(f.state,f.g,bad,mode!=2)==0,"invalid extent cannot overflow byte estimate"); } } for (int64_t upto : {0,1,3,4,5,63,64,65,96}) { const bool index=mode!=2; std::string err; f.fill(13); ConversationKv a,b,restored; cuda_check(cudaDeviceSynchronize()); check(conversation_kv_save(a,f.state,f.g,upto,index,err),"save A"); check(a.bytes()==conversation_kv_bytes(f.state,f.g,upto,index),"size estimate equals snapshot payload"); f.fill(177); check(conversation_kv_save(b,f.state,f.g,upto,index,err),"save B"); check(upto==0 || !equal(a,b),"fixture changes state"); check(conversation_kv_restore(a,f.state,f.g,upto,index,err),"restore A over B"); cuda_check(cudaDeviceSynchronize()); check(conversation_kv_save(restored,f.state,f.g,upto,index,err),"read back restored A"); check(equal(a,restored),"A/B/A byte-exact K/V and indexer state"); uint64_t fingerprint = 0; check(conversation_kv_verify(a,f.state,f.g,upto,index,fingerprint,err),"verify restored authoritative and resident bytes without rewriting them"); if (mode==1) { std::vector table((size_t)f.state.n_pages); cuda_check(cudaMemcpy(table.data(),f.state.page_table,table.size()*4,cudaMemcpyDeviceToHost)); check(std::all_of(table.begin(),table.end(),[](int32_t x){return x==-1;}),"stale VRAM pages invalidated"); } if (mode==2 && upto>0) { const auto s=qsa_real_shapes(); const int64_t b1=(upto+s.page_size-1)/s.page_size,b0=std::max(0,b1-f.state.n_slots); const size_t block=kv_block_bytes(s,fmt)/2; // Codes and scales are separate in INT8; check code bytes here. const size_t code_block=fmt==kKvInt8 ? s.page_size*s.n_head_kv*s.head_dim : block; const void* slots=fmt==kKvQ4?(void*)f.state.k_q4:fmt==kKvInt8?(void*)f.state.k_q:(void*)f.state.k_pool; std::vector got(code_block); for (int64_t page=b0;page expected(code_block); check(a.k.read(expected.data(), page*code_block, code_block),"read segmented expected draft page"); check(got == expected,"draft ring contains restored pages"); } } auto bad=a; bad.head_dim++; check(!conversation_kv_restore(bad,f.state,f.g,upto,index,err),"reject incompatible geometry"); if (fmt==kKvInt8) { // #293: INT8 K/V saved through the Hadamard rotation (STRATA_KV_ROT=1) never restores into a state // without it, nor the reverse f.state.kv_rot=true; ConversationKv rotated; check(conversation_kv_save(rotated,f.state,f.g,upto,index,err),"save rotated INT8"); check(!conversation_kv_restore(a,f.state,f.g,upto,index,err),"unrotated image into a rotated state rejected"); f.state.kv_rot=false; check(!conversation_kv_restore(rotated,f.state,f.g,upto,index,err),"rotated image into an unrotated state rejected"); check(conversation_kv_restore(a,f.state,f.g,upto,index,err),"unrotated image still restores"); } if (!a.k.empty()) { bad=a; bad.k.pop_back(); check(!conversation_kv_restore(bad,f.state,f.g,upto,index,err),"reject malformed payload"); } } } for (int fmt : {kKvF16,kKvInt8,kKvQ4}) for (int mode : {0,1}) for (int experts : {256,512}) full_session(fmt,mode,experts); for (int experts : {256,512}) full_session(3,0,experts); std::printf("conversation_snapshot_test: %d checks passed\n",checks); }