Winery-Strata / src /core /conversation_snapshot_test.cpp
Zandy-Wandy's picture
🍷 Winery Strata: fork of Niko1221/Strata (Winery Cuvée family, wine-red UI)
bbb6388 verified
Raw History Blame Contribute Delete
16.7 kB
#include "strata/core/conversation_snapshot.hpp"
#include "strata/kernels/kv_q4.hpp"
#include <cuda_runtime.h>
#include <array>
#include <cstdio>
#include <cstdlib>
#include <limits>
#include <vector>
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);
}