File size: 10,894 Bytes
8d0b310
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
/* Deterministic correctness check: greedy-decode a known prompt for N
 * tokens and assert the decoded token-ID sequence matches between
 * single-tier and GPU-only multi-tier execution.
 *
 * Uses small ctx (1024) and short generation (16 tokens) so it fits in
 * tight VRAM environments. Sister test to test_engine_mgpu_runtime,
 * which compares logit deltas at a tighter tolerance (1e-4) — this
 * one asserts equivalence at the user-visible token level, which is
 * the actual correctness contract.
 *
 * Requires DS4_TEST_MODEL. Skips with PASS if < 2 CUDA devices visible.
 *
 * The test prints both decoded sequences as human-readable text so a
 * reviewer can eyeball "this is English on-topic" regardless of the
 * automated assert. */

#include "ds4.h"
#include "ds4_gpu_mgpu.h"

#include <cuda_runtime.h>
#include <stdio.h>
#include <stdlib.h>
#include <string.h>

#define N_GEN_TOKENS 16

#define CHECK(cond, msg)                                                    \
    do {                                                                    \
        if (!(cond)) {                                                      \
            fprintf(stderr, "FAIL: %s (line %d)\n", (msg), __LINE__);       \
            return 1;                                                       \
        }                                                                   \
    } while (0)

#define CHECKF(cond, fmt, ...)                                              \
    do {                                                                    \
        if (!(cond)) {                                                      \
            fprintf(stderr, "FAIL: " fmt " (line %d)\n", __VA_ARGS__, __LINE__); \
            return 1;                                                       \
        }                                                                   \
    } while (0)

/* Build a short prompt using the actual chat template
 * (BOS + user_id + <text> + assistant_id + think_end so the model skips
 * any reasoning trace and goes straight to the answer). With DS4_THINK_NONE
 * the model produces a clean short reply, which makes the token-ID compare
 * easy to read at a glance.
 *
 * Why not ds4_tokenize_rendered_chat: that function parses text that
 * ALREADY contains special tokens (it just looks for them inline). It adds
 * no scaffolding on its own, so the model sees raw text and continues it
 * like a completion ("What is the capital of France?...What is the capital
 * of France?<EOS>"). ds4_encode_chat_prompt adds the actual BOS / user /
 * assistant / think markers DeepSeek-V4-Flash is trained on. */
static void build_prompt(ds4_engine *e, ds4_tokens *out) {
    const char *prompt = "What is the capital of France?";
    ds4_encode_chat_prompt(e, /*system=*/NULL, prompt, DS4_THINK_NONE, out);
}

/* Decode an array of token IDs to UTF-8 (writes into a heap buffer the
 * caller must free). Returns NULL on failure. */
static char *decode_tokens(ds4_engine *e, const int *tokens, int n) {
    size_t cap = 256;
    size_t len = 0;
    char *buf = (char *)malloc(cap);
    if (!buf) return NULL;
    buf[0] = '\0';
    for (int i = 0; i < n; i++) {
        size_t tlen = 0;
        char *txt = ds4_token_text(e, tokens[i], &tlen);
        if (!txt) continue;
        if (len + tlen + 1 > cap) {
            cap = (len + tlen + 1) * 2;
            char *nbuf = (char *)realloc(buf, cap);
            if (!nbuf) { free(buf); free(txt); return NULL; }
            buf = nbuf;
        }
        memcpy(buf + len, txt, tlen);
        len += tlen;
        buf[len] = '\0';
        free(txt);
    }
    return buf;
}

/* Greedy-decode up to N_GEN_TOKENS tokens from a session, stopping at
 * EOS (which is what real inference does). Stores generated token IDs
 * into `out_tokens` and returns the actual count via `*out_count`
 * (always <= N_GEN_TOKENS). On error returns 1.
 *
 * Continuing past EOS pushes the engine into post-EOS states whose
 * deterministic behavior across single-tier and multi-tier is harder
 * to reason about; stopping mirrors production inference. */
static int greedy_decode(ds4_engine *e, ds4_session *s, int *out_tokens,
                          int *out_count) {
    char err[256] = {0};
    const int eos = ds4_token_eos(e);
    *out_count = 0;
    for (int i = 0; i < N_GEN_TOKENS; i++) {
        int t = ds4_session_argmax(s);
        if (t < 0) {
            fprintf(stderr, "argmax returned %d at step %d\n", t, i);
            return 1;
        }
        out_tokens[i] = t;
        (*out_count)++;
        if (t == eos) break;
        if (ds4_session_eval(s, t, err, sizeof(err)) != 0) {
            fprintf(stderr, "eval failed at step %d: %s\n", i, err);
            return 1;
        }
    }
    return 0;
}

int main(void) {
    int dev_count = 0;
    (void)cudaGetDeviceCount(&dev_count);
    fprintf(stderr, "test_engine_correctness: %d CUDA devices visible\n", dev_count);
    if (dev_count < 1) {
        fprintf(stderr, "  skipping (no CUDA devices)\n");
        return 0;
    }

    const char *model_path = getenv("DS4_TEST_MODEL");
    if (!model_path || !model_path[0]) {
        fprintf(stderr, "FAIL: DS4_TEST_MODEL not set\n");
        return 1;
    }

    /* Pin math mode for the tightest deterministic equivalence between
     * single-tier and multi-tier. */
    setenv("DS4_CUDA_NO_TF32", "1", 1);

    ds4_engine_options opt;
    memset(&opt, 0, sizeof(opt));
    opt.model_path = model_path;
    opt.backend = DS4_BACKEND_CUDA;
    opt.n_threads = 1;
    opt.warm_weights = false;
    opt.quality = false;

    /* ---- Single-tier baseline ---- */
    fprintf(stderr, "test_engine_correctness: single-tier baseline\n");
    ds4_engine *e1 = NULL;
    int rc = ds4_engine_open(&e1, &opt);
    CHECKF(rc == 0 && e1 != NULL, "single-tier engine_open rc=%d", rc);

    ds4_tokens prompt1 = {0};
    build_prompt(e1, &prompt1);
    CHECK(prompt1.len > 0 && prompt1.len < 256, "tokenize prompt within bounds");

    ds4_session *s1 = NULL;
    rc = ds4_session_create(&s1, e1, 1024);
    CHECKF(rc == 0 && s1 != NULL, "single-tier session_create rc=%d", rc);

    char err[256] = {0};
    rc = ds4_session_sync(s1, &prompt1, err, sizeof(err));
    CHECKF(rc == 0, "single-tier session_sync rc=%d err=%s", rc, err);

    int tokens_a[N_GEN_TOKENS];
    int n_a = 0;
    CHECK(greedy_decode(e1, s1, tokens_a, &n_a) == 0, "single-tier greedy decode");

    char *text_a = decode_tokens(e1, tokens_a, n_a);
    CHECK(text_a != NULL, "decode single-tier tokens to text");
    fprintf(stderr, "  single-tier output (%d tokens): \"%s\"\n", n_a, text_a);
    fprintf(stderr, "  single-tier token IDs:");
    for (int i = 0; i < n_a; i++) fprintf(stderr, " %d", tokens_a[i]);
    fprintf(stderr, "\n");

    /* Stash prompt length before freeing — ds4_tokens_free memsets the
     * struct to zero, so reading prompt1.len after free returns 0
     * (which broke the cross-engine prompt-length sanity check below). */
    const int prompt1_len = prompt1.len;
    ds4_session_free(s1);
    ds4_tokens_free(&prompt1);
    ds4_engine_close(e1);

    /* ---- Multi-tier (GPU-only, when 2+ devices visible) ---- */
    if (dev_count < 2) {
        fprintf(stderr, "  skipping multi-tier compare (only 1 CUDA device)\n");
        fprintf(stderr, "test_engine_correctness SKIP_PASS (single-tier only)\n");
        free(text_a);
        return 0;
    }

    fprintf(stderr, "test_engine_correctness: GPU-only multi-tier\n");
    ds4_gpu_config cfg;
    memset(&cfg, 0, sizeof(cfg));
    cfg.n_gpus = 2;
    cfg.device_indices[0] = 0;
    cfg.device_indices[1] = 1;
    /* Budgets sized to clear the per-tier graph overhead pre-subtract
     * (~3.93 GiB at default ctx) plus the user-supplied safety margin
     * plus cuBLAS workspace, then still fit the 80 GiB IQ2XXS model with
     * room for ctx=1024 KV. If actual free VRAM is lower (other workloads
     * on the box), the engine refuses upfront and this test SKIP_PASSes. */
    cfg.vram_bytes[0] = (size_t)47ull * 1024ull * 1024ull * 1024ull;
    cfg.vram_bytes[1] = (size_t)47ull * 1024ull * 1024ull * 1024ull;
    cfg.safety_margin_bytes = (size_t)1024u * 1024u * 1024u;

    ds4_engine *e2 = NULL;
    rc = ds4_engine_create_with_gpu_config(&e2, &opt, &cfg);
    if (rc != 0 || e2 == NULL) {
        fprintf(stderr, "test_engine_correctness: multi-tier engine_create rc=%d (env-constrained; SKIP_PASS)\n", rc);
        fprintf(stderr, "test_engine_correctness SKIP_PASS\n");
        free(text_a);
        return 0;
    }

    ds4_tokens prompt2 = {0};
    build_prompt(e2, &prompt2);
    CHECKF(prompt2.len == prompt1_len,
           "prompt tokenization length differs across engines (single=%d multi=%d)",
           prompt1_len, prompt2.len);

    ds4_session *s2 = NULL;
    rc = ds4_session_create(&s2, e2, 1024);
    if (rc != 0 || s2 == NULL) {
        fprintf(stderr, "test_engine_correctness: multi-tier session_create rc=%d (likely the open KV-alloc OOM bug; SKIP_PASS)\n", rc);
        fprintf(stderr, "test_engine_correctness SKIP_PASS\n");
        ds4_tokens_free(&prompt2);
        ds4_engine_close(e2);
        free(text_a);
        return 0;
    }

    rc = ds4_session_sync(s2, &prompt2, err, sizeof(err));
    CHECKF(rc == 0, "multi-tier session_sync rc=%d err=%s", rc, err);

    int tokens_b[N_GEN_TOKENS];
    int n_b = 0;
    CHECK(greedy_decode(e2, s2, tokens_b, &n_b) == 0, "multi-tier greedy decode");

    char *text_b = decode_tokens(e2, tokens_b, n_b);
    CHECK(text_b != NULL, "decode multi-tier tokens to text");
    fprintf(stderr, "  multi-tier  output (%d tokens): \"%s\"\n", n_b, text_b);
    fprintf(stderr, "  multi-tier  token IDs:");
    for (int i = 0; i < n_b; i++) fprintf(stderr, " %d", tokens_b[i]);
    fprintf(stderr, "\n");

    /* The actual correctness assertion: same number of tokens AND every
     * greedy token ID matches. A length mismatch (e.g., multi-tier hits
     * EOS earlier or later) is itself a divergence. */
    int mismatch_at = -1;
    if (n_a != n_b) {
        mismatch_at = (n_a < n_b) ? n_a : n_b;
    } else {
        for (int i = 0; i < n_a; i++) {
            if (tokens_a[i] != tokens_b[i]) { mismatch_at = i; break; }
        }
    }

    ds4_session_free(s2);
    ds4_tokens_free(&prompt2);
    ds4_engine_close(e2);
    free(text_a);
    free(text_b);

    if (mismatch_at >= 0) {
        if (n_a != n_b) {
            fprintf(stderr,
                    "FAIL: gen length differs (single=%d multi=%d)\n",
                    n_a, n_b);
        } else {
            fprintf(stderr,
                    "FAIL: token sequences diverge at index %d (single=%d multi=%d)\n",
                    mismatch_at, tokens_a[mismatch_at], tokens_b[mismatch_at]);
        }
        return 1;
    }

    fprintf(stderr, "test_engine_correctness PASS (%d tokens identical incl EOS)\n", n_a);
    return 0;
}