File size: 14,930 Bytes
bbb6388 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 | // Host-only validation fixtures. CUDA is linked but must never be initialized:
// invalid restores must return before the first CUDA call or destination write.
#include "strata/core/conversation_snapshot.hpp"
#include "strata/kernels/kv_q4.hpp"
#include <cstdio>
#include <cstdlib>
#include <functional>
#include <limits>
#include <cstring>
#if defined(CONVERSATION_TEST_TRANSFERS)
// GNU/ELF link wrapping exercises the actual restore control flow without
// initializing CUDA or deliberately poisoning a real device context. This is
// host-backend fault injection, not a claim of hardware-failure recovery.
namespace {
int copy_calls = 0, sync_calls = 0, fail_copy = 0, fail_sync = 0;
size_t copied_bytes = 0;
}
extern "C" cudaError_t __wrap_cudaMemcpy(void* dst, const void* src, size_t n, cudaMemcpyKind) {
if (++copy_calls == fail_copy) return cudaErrorInvalidValue;
copied_bytes += n;
std::memcpy(dst, src, n);
return cudaSuccess;
}
extern "C" cudaError_t __wrap_cudaDeviceSynchronize() {
return ++sync_calls == fail_sync ? cudaErrorUnknown : cudaSuccess;
}
extern "C" cudaError_t __wrap_cudaGetLastError() { return cudaSuccess; }
#endif
using namespace strata::core;
namespace {
int checks = 0;
void check(bool ok, const char* label) {
++checks;
if (!ok) { std::fprintf(stderr, "FAIL: %s\n", label); std::exit(1); }
}
struct Pools {
QsaState st;
std::array<std::vector<uint8_t>, 8> data;
Pools(const ModelGeometry& g, int format) {
st.max_cells = 96; st.n_pages = st.n_slots = 24; st.idx_pooled_rows = 26;
st.kv_int8 = format == 1; st.kv_q4 = format == 2;
const size_t per = format == 2 ? strata::kernels::kv_q4_bytes_per_head((int) g.head_dim)
: g.head_dim * (format == 1 ? 1 : 2);
data[0].resize(96 * g.n_head_kv * per, 0xa5); data[1] = data[0];
data[2].resize(format == 1 ? 96 * g.n_head_kv * (g.head_dim / 64) * 2 : 0, 0xa5); data[3] = data[2];
data[4].resize(26 * g.idx_key_dim * 4, 0xa5);
data[5].resize(3 * g.idx_key_dim * 4, 0xa5);
data[6].resize(g.idx_key_dim * 4, 0xa5); data[7].resize(4, 0xa5);
if (format == 2) { st.k_q4 = data[0].data(); st.v_q4 = data[1].data(); }
else if (format == 1) {
st.k_q = (int8_t*) data[0].data(); st.v_q = (int8_t*) data[1].data();
st.k_scale = (uint16_t*) data[2].data(); st.v_scale = (uint16_t*) data[3].data();
} else { st.k_pool = (uint16_t*) data[0].data(); st.v_pool = (uint16_t*) data[1].data(); }
st.idx_pooled = (float*) data[4].data(); st.idx_tail = (float*) data[5].data();
st.idx_dead = (float*) data[6].data(); st.idx_block_pos = (int32_t*) data[7].data();
}
ConversationKv image(const ModelGeometry& g, bool index) const {
ConversationKv k;
k.format = qsa_kv_format(st); k.cells = 12; k.page_size = 4;
k.heads = g.n_head_kv; k.head_dim = g.head_dim; k.idx_dim = g.idx_key_dim;
k.pooled_rows = index ? 3 : 0;
k.k.resize(data[0].size() / 8, 13); k.v = k.k;
k.k_scale.resize(data[2].size() / 8, 13); k.v_scale = k.k_scale;
k.pooled.resize(k.pooled_rows * g.idx_key_dim * 4, 13);
return k;
}
};
void fixture(int format, int experts, bool zero_qsa, bool ple) {
ModelGeometry g;
g.n_layers = zero_qsa ? 3 : 8; g.n_expert = experts;
g.ssm_state_size = 2; g.ssm_v_heads = 2; g.ssm_conv_channels = 8;
g.n_head_kv = 1; g.head_dim = 64; g.idx_key_dim = 8;
std::string error;
ConversationStateSizes z;
check(conversation_state_sizes(g, z, error), "checked canonical/pruned/zero-QSA sizing");
Pools first(g, format), last(g, format), draft(g, format);
std::array<QsaState, 2> layers{first.st, last.st};
std::vector<uint8_t> gdn(z.gdn, 0xa5), history(ple ? z.ple : 0, 0xa5);
SessionState ss;
ss.max_cells = 96; ss.gdn_state = (float*) gdn.data();
// the whole-model carve, as session_init leaves a one-GPU session
ss.layer_hi = g.n_layers; ss.gdn_alloc = g.n_gdn_layers(); ss.qsa_alloc = g.n_qsa_layers();
ss.ple_hist = ple ? (float*) history.data() : nullptr;
ss.qsa_states = zero_qsa ? nullptr : layers.data();
SavedConversation image;
image.geometry = {g.n_embd,g.n_layers,g.qsa_interval,g.ssm_state_size,g.ssm_k_heads,g.ssm_v_heads,
g.ssm_d_conv,g.ssm_conv_channels,g.ssm_value_dim,g.n_head,g.n_head_kv,g.head_dim,
g.idx_q_heads,g.idx_key_dim,g.hc,g.hc_lr,g.n_expert,g.n_ff};
image.layer_hi = g.n_layers;
auto checkpoint = [&](size_t tokens) {
ConversationCheckpoint c;
c.ids.resize(tokens);
for (size_t i = 0; i < tokens; ++i) c.ids[i] = (int32_t) i + 1;
c.imgs = {{1, 44}};
c.gdn.resize(z.gdn, 13); c.ple.resize(ple ? z.ple : 0, 13);
c.tails.resize(g.n_qsa_layers() * z.tail, 13); c.dead.resize(g.n_qsa_layers() * z.dead, 13);
c.block_pos.resize(g.n_qsa_layers() * z.block_pos, 13);
return c;
};
image.live = checkpoint(9); image.checkpoints.push_back(checkpoint(5));
image.kv.reserve((size_t) g.n_qsa_layers() + 1);
if (!zero_qsa) { image.kv.push_back(first.image(g, true)); image.kv.push_back(last.image(g, true)); }
image.kv.push_back(draft.image(g, false));
check(conversation_snapshot_validate(image, ss, g, draft.st, error), "complete image validates without CUDA");
size_t estimate = 0;
check(conversation_snapshot_bytes({image.live.ids,image.live.imgs,image.checkpoints,true},ss,g,draft.st,estimate,error),
"capture estimate works without CUDA");
check(estimate == image.bytes(), "estimate covers checkpoint/indexer/spare payloads");
auto unchanged = [&] {
auto pristine = [](const auto& bytes) { return std::all_of(bytes.begin(),bytes.end(),[](uint8_t b){return b==0xa5;}); };
if (!pristine(gdn) || !pristine(history) || ss.ple_prev[0] != -1 || ss.ple_prev[1] != -1) return false;
for (const auto* p : {&first,&last,&draft}) for (const auto& bytes : p->data) if (!pristine(bytes)) return false;
return true;
};
auto reject = [&](const std::function<void(SavedConversation&)>& mutate, const char* label) {
auto bad = image; mutate(bad);
check(conversation_snapshot_restore(bad,ss,g,draft.st,error) == ConversationRestore::invalid,label);
check(unchanged(), "invalid restore did not touch any layer or running-state buffer");
};
reject([](auto& s){s.kv.back().k.pop_back();}, "late draft corruption rejected before first layer write");
reject([](auto& s){s.geometry[16] = 128;}, "different expert geometry rejected");
reject([](auto& s){s.live.stage_parts.emplace_back();}, "layer-split live state rejected before writes");
reject([](auto& s){s.checkpoints[0].stage_parts.emplace_back();}, "layer-split checkpoint rejected before writes");
reject([](auto& s){s.kv.pop_back();}, "missing KV layer rejected");
reject([](auto& s){s.live.gdn.pop_back();}, "bad live recurrence rejected");
reject([](auto& s){s.checkpoints.back().gdn.pop_back();}, "bad retained checkpoint rejected");
reject([](auto& s){s.checkpoints.back().ids[0] = 99;}, "foreign checkpoint prefix rejected");
reject([](auto& s){s.checkpoints.back().imgs[0].hash++;}, "foreign checkpoint image rejected");
reject([](auto& s){s.live.imgs[0].start = -1;}, "negative image position rejected");
reject([](auto& s){s.live.ids[0] = -1;}, "negative live token rejected");
reject([](auto& s){s.live.ple.push_back(0);}, "PLE size/enabled mismatch rejected");
reject([](auto& s){s.live.tails.push_back(0);}, "tail size/zero-QSA mismatch rejected");
if (!zero_qsa) {
reject([](auto& s){s.kv[1].pooled.pop_back();}, "late indexer spare row corruption rejected");
reject([](auto& s){s.live.dead.pop_back();}, "missing indexer spare key rejected");
reject([](auto& s){s.live.block_pos.pop_back();}, "missing indexer metadata rejected");
ss.qsa_states = nullptr;
check(!conversation_snapshot_validate(image,ss,g,draft.st,error), "missing QSA targets rejected");
ss.qsa_states = layers.data();
reject([](auto& s){s.layer_lo = 4;}, "snapshot from another layer range rejected");
// #216's carve: a split stage owning layers [4, 8) holds GDN rows 3..5 and QSA ordinal 1 only
SessionState stage = ss;
stage.layer_lo = 4; stage.gdn_alloc = 3; stage.qsa_ord0 = 1; stage.qsa_alloc = 1;
ConversationStateSizes zs;
check(conversation_session_sizes(g,stage,zs,error) && zs.gdn == z.gdn / 6 * 3 && zs.tail == z.tail,
"carved session sizes cover the owned rows only");
check(!conversation_checkpoint_validate(image.live,stage,g,error), "whole-model checkpoint rejected by a carve");
ConversationCheckpoint part = image.live;
part.gdn.resize(zs.gdn); part.tails.resize(zs.tail); part.dead.resize(zs.dead);
part.block_pos.resize(zs.block_pos);
check(conversation_checkpoint_validate(part,stage,g,error), "carve-sized checkpoint validates");
check(!conversation_snapshot_validate(image,stage,g,draft.st,error), "whole-model snapshot rejected by a carve");
stage.qsa_alloc = 2;
check(!conversation_session_sizes(g,stage,zs,error), "carve past the last QSA ordinal rejected");
stage.qsa_alloc = 1; stage.gdn_alloc = 7;
check(!conversation_session_sizes(g,stage,zs,error), "carve past the GDN rows rejected");
}
auto bad_geometry = g; bad_geometry.ssm_state_size = std::numeric_limits<int64_t>::max();
check(!conversation_state_sizes(bad_geometry,z,error), "running-state arithmetic overflow rejected");
bad_geometry = g; bad_geometry.qsa_interval = 0;
check(!conversation_state_sizes(bad_geometry,z,error), "zero layer interval rejected before division");
bad_geometry = g; bad_geometry.n_head_kv = std::numeric_limits<int64_t>::max();
check(conversation_kv_bytes(draft.st,bad_geometry,9,false)==0, "KV byte overflow rejected");
#if defined(CONVERSATION_TEST_TRANSFERS)
auto reset = [&] {
for (auto* p : {&first,&last,&draft}) for (auto& bytes : p->data) std::fill(bytes.begin(),bytes.end(),0xa5);
std::fill(gdn.begin(),gdn.end(),0xa5); std::fill(history.begin(),history.end(),0xa5);
ss.ple_prev[0] = ss.ple_prev[1] = -1;
copy_calls = sync_calls = fail_copy = fail_sync = 0;
};
reset(); fail_sync = 1;
check(conversation_snapshot_restore(image,ss,g,draft.st,error)==ConversationRestore::transfer_failed,
"pre-transfer CUDA synchronization failure is fatal");
check(copy_calls==0 && unchanged(), "pre-sync failure did not mutate state");
reset(); fail_copy = 1;
check(conversation_snapshot_restore(image,ss,g,draft.st,error)==ConversationRestore::transfer_failed,
"first CUDA transfer failure is fatal");
check(unchanged(), "first-copy failure did not mutate state");
reset(); fail_copy = 2;
check(conversation_snapshot_restore(image,ss,g,draft.st,error)==ConversationRestore::transfer_failed,
"partial CUDA transfer failure is fatal, not an invalid-image fallback");
check(!unchanged(), "fault fixture genuinely produced partial state");
reset(); fail_sync = 2;
check(conversation_snapshot_restore(image,ss,g,draft.st,error)==ConversationRestore::transfer_failed,
"post-transfer CUDA synchronization failure is fatal");
check(copy_calls>0 && !unchanged(), "post-sync failure occurred after state writes");
reset();
check(conversation_snapshot_restore(image,ss,g,draft.st,error)==ConversationRestore::restored,
"injected host transfer backend completes a valid restore");
check(gdn==image.live.gdn && history==image.live.ple && ss.ple_prev[0]==8 && ss.ple_prev[1]==9,
"successful restore publishes correct running state and PLE window");
uint64_t fingerprint = 0;
check(conversation_kv_verify(image.kv.back(),draft.st,g,9,false,fingerprint,error), "read-back verifies complete draft payload");
draft.data[0][0] ^= 1;
check(!conversation_kv_verify(image.kv.back(),draft.st,g,9,false,fingerprint,error), "read-back detects corrupted draft byte");
draft.data[0][0] ^= 1;
fail_copy = copy_calls + 1;
check(!conversation_kv_verify(image.kv.back(),draft.st,g,9,false,fingerprint,error), "read-back transfer failure is not a successful fingerprint");
reset(); copied_bytes = 0;
const ConversationView view{image.live.ids,image.live.imgs,image.checkpoints,true};
SavedConversation full,incremental;
check(conversation_snapshot_save(full,view,ss,g,draft.st,error), "full capture through host transfer backend");
const size_t full_copies = copied_bytes;
ConversationKvReuse reuse{full.kv,9,9};
size_t peak = 0, reused = 0;
check(conversation_snapshot_capture_bytes(reuse,view,ss,g,draft.st,peak,error), "incremental capture admission without transfers");
copied_bytes = 0;
const bool saved_incrementally = conversation_snapshot_save(incremental,view,ss,g,draft.st,error,std::move(reuse),&reused);
if (!saved_incrementally) std::fprintf(stderr,"capture error: %s\n",error.c_str());
check(saved_incrementally, "incremental capture through host backend");
check(reused > 0 && copied_bytes + reused == full_copies, "reused byte count measures transfers actually omitted");
check(incremental.bytes() <= peak && incremental.live.gdn == full.live.gdn, "running state refreshed within admitted allocation");
for (size_t i=0;i<full.kv.size();++i)
check(incremental.kv[i].k==full.kv[i].k && incremental.kv[i].v==full.kv[i].v &&
incremental.kv[i].k_scale==full.kv[i].k_scale && incremental.kv[i].v_scale==full.kv[i].v_scale &&
incremental.kv[i].pooled==full.kv[i].pooled, "incremental and full capture agree across every payload");
reuse = {full.kv,9,9};
reuse.kv.back().k.pop_back();
copy_calls = sync_calls = 0;
check(!conversation_snapshot_capture_bytes(reuse,view,ss,g,draft.st,peak,error), "malformed retained draft rejected before admission");
check(!copy_calls && !sync_calls, "invalid retained storage performs no CUDA calls");
reuse = {full.kv,9,9};
fail_copy = copy_calls + 1;
SavedConversation unpublished;
check(!conversation_snapshot_save(unpublished,view,ss,g,draft.st,error,std::move(reuse)), "incremental capture copy failure is reported");
check(unpublished.kv.empty() && unpublished.live.ids.empty(), "failed incremental capture cannot publish a partial image");
#endif
}
}
int main() {
for (int format : {0,1,2}) for (int experts : {256,512})
for (bool zero_qsa : {false,true}) for (bool ple : {false,true}) fixture(format,experts,zero_qsa,ple);
std::printf("conversation_validation_test: %d host-only checks passed\n", checks);
}
|