Download src/core/conversation_snapshot_test.cpp from WineryLabs/Winery-Strata: direct link, hf CLI and curl.
- Browser
- Download file 16.7 kB
-
https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/core/conversation_snapshot_test.cpp
- Command line
-
hf download hf://WineryLabs/Winery-Strata/src/core/conversation_snapshot_test.cpp
-
curl -L -o conversation_snapshot_test.cpp https://huggingface.co/WineryLabs/Winery-Strata/resolve/main/src/core/conversation_snapshot_test.cpp
16.7 kB
| 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<void*> device, host; | |
| std::array<void*,5> sources{}; | |
| std::array<size_t,5> sizes{}; | |
| template<class T> 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<T*>(mapped); | |
| } else { | |
| cuda_check(cudaMalloc(&raw, n)); device.push_back(raw); p = static_cast<T*>(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<sources.size();++i) { | |
| std::vector<uint8_t> data(sizes[i]); | |
| for (size_t j=0;j<data.size();++j) data[j]=(uint8_t)(salt+i*31+j*7+j/257); | |
| if (!data.empty()) cuda_check(cudaMemcpy(sources[i],data.data(),data.size(),cudaMemcpyDefault)); | |
| } | |
| } | |
| void fill_after(uint8_t salt, int64_t first_dirty) { | |
| for (size_t i=0;i<sources.size();++i) { | |
| const size_t offset = i == 4 ? size_t(first_dirty/4)*g.idx_key_dim*4 | |
| : (sizes[i]/size_t(state.max_cells))*size_t((first_dirty/4)*4); | |
| if (sizes[i] > offset) | |
| cuda_check(cudaMemset(static_cast<uint8_t*>(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<int32_t> ids(65); | |
| for (size_t i=0;i<ids.size();++i) ids[i]=(int32_t)i+1; | |
| std::vector<ConversationImageKey> images; | |
| std::vector<ConversationCheckpoint> checkpoints; | |
| auto fill=[&](uint8_t salt) { | |
| main.fill(salt); draft.fill(salt); | |
| for (const auto& [p,n] : std::vector<std::pair<void*,size_t>>{ | |
| {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<int64_t>(0,dirty-1)); | |
| cuda_check(cudaMemset(ss.gdn_state,91,sizes.gdn)); | |
| ids.resize(70); | |
| for (size_t i=65;i<ids.size();++i) ids[i]=int32_t(i+1); | |
| cuda_check(cudaMemcpy(main.state.idx_pooled+(ids.size()/4)*g.idx_key_dim, | |
| main.state.idx_dead,sizes.dead,cudaMemcpyDeviceToDevice)); | |
| SavedConversation fresh,incremental; | |
| check(conversation_snapshot_save(fresh,view,ss,g,draft.state,err),"full capture reference after growth or rewind"); | |
| size_t peak=0,reused=0; | |
| check(conversation_snapshot_capture_bytes(reuse,view,ss,g,draft.state,peak,err),"admit incremental capture peak"); | |
| check(conversation_snapshot_save(incremental,view,ss,g,draft.state,err,std::move(reuse),&reused),"capture with retained pages"); | |
| check(incremental.bytes() <= peak,"incremental allocation stays within admitted bound"); | |
| check(incremental.live.gdn==fresh.live.gdn && incremental.live.ple==fresh.live.ple && | |
| incremental.live.dead==fresh.live.dead && equal(incremental.kv[0],fresh.kv[0]) && | |
| equal(incremental.kv[1],fresh.kv[1]),"incremental capture equals full capture after growth or rewind"); | |
| check((dirty>=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<uint8_t> 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<int,4>{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<int64_t>::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<int32_t> 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<int64_t>(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<uint8_t> got(code_block); | |
| for (int64_t page=b0;page<b1;++page) { | |
| cuda_check(cudaMemcpy(got.data(),(const uint8_t*)slots+(page%f.state.n_slots)*code_block,code_block,cudaMemcpyDeviceToHost)); | |
| std::vector<uint8_t> 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); | |
| } | |