| #ifndef LAL_DATA_LOADER_H |
| #define LAL_DATA_LOADER_H |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| #include <stdint.h> |
| #include <stdio.h> |
| #include <stdlib.h> |
| #include <string.h> |
|
|
| #define LAL_DATA_MAGIC "LALT" |
| #define LAL_DATA_MAGIC_LEN 4 |
| #define LAL_DATA16_MAGIC "LALT16" |
| #define LAL_DATA16_MAGIC_LEN 6 |
|
|
| |
| typedef struct { |
| int n_tokens; |
| int *tokens; |
| } TrainSample; |
|
|
| |
| |
| |
| |
| #define PREFETCH_BATCH 8 |
| typedef struct { |
| int idx; |
| int n_tokens; |
| int tokens[8192]; |
| } PrefetchSlot; |
|
|
| typedef struct { |
| PrefetchSlot slots[PREFETCH_BATCH]; |
| int head; |
| int tail; |
| int count; |
| int active; |
| } PrefetchBuffer; |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| typedef struct { |
| FILE *f; |
| int n_samples; |
| int64_t *offsets; |
| int max_len; |
| int n_vocab; |
| int tok_width; |
| |
| int pf_next_idx; |
| int pf_cache_idx[PREFETCH_BATCH]; |
| int pf_cache_len[PREFETCH_BATCH]; |
| int *pf_cache_data; |
| int pf_head; |
| int pf_count; |
| } DataLoader; |
|
|
| |
|
|
| |
| |
| |
| |
| int dataloader_init(DataLoader *dl, const char *path); |
|
|
| |
| |
| |
| int dataloader_get(DataLoader *dl, int idx, int *tokens, int max_len); |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| int dataloader_random_batch(DataLoader *dl, int *indices, int batch_size, |
| int **tokens, int *lengths); |
|
|
| |
| |
| void dataloader_free_batch(int **tokens, int batch_size); |
|
|
| |
| void dataloader_free(DataLoader *dl); |
|
|
| |
| |
| |
| |
| |
| |
| |
| #if defined(LAL_DATA_LOADER_IMPLEMENTATION) || defined(_LAL_DATA_LOADER_TEST) |
|
|
| #include <errno.h> |
|
|
| int dataloader_init(DataLoader *dl, const char *path) { |
| if (!dl || !path) return 1; |
| memset(dl, 0, sizeof(*dl)); |
|
|
| FILE *f = fopen(path, "rb"); |
| if (!f) { |
| fprintf(stderr, "[dataloader] cannot open '%s': %s\n", |
| path, strerror(errno)); |
| return 2; |
| } |
|
|
| |
| |
| |
| char magic[LAL_DATA16_MAGIC_LEN]; |
| if (fread(magic, 1, LAL_DATA16_MAGIC_LEN, f) != LAL_DATA16_MAGIC_LEN) { |
| fprintf(stderr, "[dataloader] truncated header in '%s'\n", path); |
| fclose(f); |
| return 3; |
| } |
| int tok_width; |
| if (memcmp(magic, LAL_DATA16_MAGIC, LAL_DATA16_MAGIC_LEN) == 0) { |
| tok_width = 2; |
| } else if (memcmp(magic, LAL_DATA_MAGIC, LAL_DATA_MAGIC_LEN) == 0) { |
| tok_width = 4; |
| if (fseek(f, LAL_DATA_MAGIC_LEN, SEEK_SET) != 0) { fclose(f); return 3; } |
| } else { |
| fprintf(stderr, "[dataloader] bad magic in '%s' (expected \"LALT\"/\"LALT16\")\n", path); |
| fclose(f); |
| return 3; |
| } |
|
|
| int32_t n_samples = 0, n_vocab = 0; |
| if (fread(&n_samples, sizeof(int32_t), 1, f) != 1 || |
| fread(&n_vocab, sizeof(int32_t), 1, f) != 1) { |
| fprintf(stderr, "[dataloader] truncated header in '%s'\n", path); |
| fclose(f); |
| return 4; |
| } |
|
|
| dl->f = f; |
| dl->n_samples = (int)n_samples; |
| dl->n_vocab = (int)n_vocab; |
| dl->max_len = 512; |
| dl->offsets = NULL; |
| dl->tok_width = tok_width; |
|
|
| |
| dl->pf_next_idx = 0; |
| dl->pf_head = 0; |
| dl->pf_count = 0; |
| dl->pf_cache_data = NULL; |
| for (int i = 0; i < PREFETCH_BATCH; i++) { |
| dl->pf_cache_idx[i] = -1; |
| dl->pf_cache_len[i] = 0; |
| } |
|
|
| if (dl->n_samples <= 0) { |
| |
| return 0; |
| } |
|
|
| |
| dl->offsets = (int64_t *)malloc(sizeof(int64_t) * (size_t)dl->n_samples); |
| if (!dl->offsets) { |
| fprintf(stderr, "[dataloader] OOM allocating offsets (%d)\n", dl->n_samples); |
| fclose(f); |
| dl->f = NULL; |
| return 5; |
| } |
|
|
| long cur = ftell(f); |
| for (int i = 0; i < dl->n_samples; i++) { |
| dl->offsets[i] = (int64_t)cur; |
| int32_t n_tok = 0; |
| if (fread(&n_tok, sizeof(int32_t), 1, f) != 1) { |
| fprintf(stderr, "[dataloader] truncated at sample %d (reading n_tokens)\n", i); |
| free(dl->offsets); |
| dl->offsets = NULL; |
| fclose(f); |
| dl->f = NULL; |
| return 6; |
| } |
| if (n_tok < 0) { |
| fprintf(stderr, "[dataloader] negative n_tokens=%d at sample %d\n", n_tok, i); |
| free(dl->offsets); |
| dl->offsets = NULL; |
| fclose(f); |
| dl->f = NULL; |
| return 7; |
| } |
| |
| if (n_tok > 0) { |
| if (fseek(f, (long)dl->tok_width * n_tok, SEEK_CUR) != 0) { |
| fprintf(stderr, "[dataloader] seek failed at sample %d\n", i); |
| free(dl->offsets); |
| dl->offsets = NULL; |
| fclose(f); |
| dl->f = NULL; |
| return 8; |
| } |
| } |
| cur += (long)sizeof(int32_t) + (long)dl->tok_width * n_tok; |
| } |
|
|
| return 0; |
| } |
|
|
| int dataloader_get(DataLoader *dl, int idx, int *tokens, int max_len) { |
| if (!dl || !tokens) return -1; |
| if (!dl->f || !dl->offsets) return -2; |
| if (idx < 0 || idx >= dl->n_samples) return -3; |
| if (max_len <= 0) return -4; |
|
|
| |
| if (dl->pf_cache_data) { |
| for (int i = 0; i < PREFETCH_BATCH; i++) { |
| int slot = (dl->pf_head + i) % PREFETCH_BATCH; |
| if (dl->pf_cache_idx[slot] == idx) { |
| |
| int n = dl->pf_cache_len[slot]; |
| int to_copy = (n > max_len) ? max_len : n; |
| memcpy(tokens, &dl->pf_cache_data[(size_t)slot * 8192], |
| sizeof(int) * to_copy); |
| if (max_len > dl->max_len) dl->max_len = max_len; |
| return to_copy; |
| } |
| } |
| } |
|
|
| |
| if (fseek(dl->f, (long)dl->offsets[idx], SEEK_SET) != 0) return -5; |
|
|
| int32_t n_tok = 0; |
| if (fread(&n_tok, sizeof(int32_t), 1, dl->f) != 1) return -6; |
|
|
| int to_read = (n_tok > max_len) ? max_len : (int)n_tok; |
| if (to_read > 0) { |
| if (dl->tok_width == 2) { |
| |
| uint16_t u16buf[8192]; |
| int done = 0; |
| while (done < to_read) { |
| int chunk = to_read - done; |
| if (chunk > 8192) chunk = 8192; |
| if (fread(u16buf, sizeof(uint16_t), (size_t)chunk, dl->f) |
| != (size_t)chunk) { |
| return -7; |
| } |
| for (int k = 0; k < chunk; k++) |
| tokens[done + k] = (int)u16buf[k]; |
| done += chunk; |
| } |
| } else { |
| if (fread(tokens, sizeof(int32_t), (size_t)to_read, dl->f) |
| != (size_t)to_read) { |
| return -7; |
| } |
| } |
| } |
|
|
| |
| if (!dl->pf_cache_data && dl->max_len > 0) { |
| dl->pf_cache_data = (int *)malloc((size_t)PREFETCH_BATCH * 8192 * sizeof(int)); |
| } |
| if (dl->pf_cache_data) { |
| |
| for (int i = 0; i < PREFETCH_BATCH; i++) { |
| int rand_idx = rand() % dl->n_samples; |
| if (fseek(dl->f, (long)dl->offsets[rand_idx], SEEK_SET) != 0) break; |
| int32_t nt = 0; |
| if (fread(&nt, sizeof(int32_t), 1, dl->f) != 1) break; |
| int tr = (nt > 8192) ? 8192 : (int)nt; |
| if (tr > 0) { |
| if (dl->tok_width == 2) { |
| |
| uint16_t u16buf[8192]; |
| if (fread(u16buf, sizeof(uint16_t), (size_t)tr, dl->f) |
| != (size_t)tr) break; |
| int *dst = &dl->pf_cache_data[(size_t)i * 8192]; |
| for (int k = 0; k < tr; k++) dst[k] = (int)u16buf[k]; |
| } else { |
| if (fread(&dl->pf_cache_data[(size_t)i * 8192], sizeof(int32_t), |
| (size_t)tr, dl->f) != (size_t)tr) break; |
| } |
| } |
| dl->pf_cache_idx[i] = rand_idx; |
| dl->pf_cache_len[i] = tr; |
| } |
| dl->pf_head = 0; |
| } |
|
|
| if (max_len > dl->max_len) dl->max_len = max_len; |
| return to_read; |
| } |
|
|
| int dataloader_random_batch(DataLoader *dl, int *indices, int batch_size, |
| int **tokens, int *lengths) { |
| if (!dl || !tokens || !lengths) return -1; |
| if (batch_size <= 0) return -2; |
| if (!indices) return -3; |
|
|
| int max_len = dl->max_len > 0 ? dl->max_len : 512; |
|
|
| for (int i = 0; i < batch_size; i++) { |
| tokens[i] = (int *)malloc(sizeof(int) * (size_t)max_len); |
| if (!tokens[i]) { |
| |
| for (int j = 0; j < i; j++) { free(tokens[j]); tokens[j] = NULL; } |
| return -(10 + i); |
| } |
| int n = dataloader_get(dl, indices[i], tokens[i], max_len); |
| if (n < 0) { |
| free(tokens[i]); tokens[i] = NULL; |
| for (int j = 0; j < i; j++) { free(tokens[j]); tokens[j] = NULL; } |
| return -(100 + i); |
| } |
| lengths[i] = n; |
| } |
| return batch_size; |
| } |
|
|
| void dataloader_free_batch(int **tokens, int batch_size) { |
| if (!tokens || batch_size <= 0) return; |
| for (int i = 0; i < batch_size; i++) { |
| if (tokens[i]) { free(tokens[i]); tokens[i] = NULL; } |
| } |
| |
| } |
|
|
| void dataloader_free(DataLoader *dl) { |
| if (!dl) return; |
| if (dl->f) { fclose(dl->f); dl->f = NULL; } |
| if (dl->offsets) { free(dl->offsets); dl->offsets = NULL; } |
| if (dl->pf_cache_data) { free(dl->pf_cache_data); dl->pf_cache_data = NULL; } |
| dl->n_samples = 0; |
| dl->n_vocab = 0; |
| } |
|
|
| #endif |
|
|
| |
| |
| |
| #ifdef _LAL_DATA_LOADER_TEST |
|
|
| #include <time.h> |
|
|
| |
| static void lal_dl_shuffle(int *a, int n, unsigned int seed) { |
| srand(seed); |
| for (int i = n - 1; i > 0; i--) { |
| int j = rand() % (i + 1); |
| int t = a[i]; a[i] = a[j]; a[j] = t; |
| } |
| } |
|
|
| int main(int argc, char **argv) { |
| const char *path = (argc > 1) ? argv[1] : "data/train_tokens.bin"; |
| printf("[test] opening %s\n", path); |
|
|
| DataLoader dl; |
| int rc = dataloader_init(&dl, path); |
| if (rc != 0) { |
| printf("[test] dataloader_init failed: %d\n", rc); |
| return 1; |
| } |
| printf("[test] n_samples=%d n_vocab=%d max_len=%d\n", |
| dl.n_samples, dl.n_vocab, dl.max_len); |
|
|
| |
| int *buf = (int *)malloc(sizeof(int) * (size_t)dl.max_len); |
| int n0 = dataloader_get(&dl, 0, buf, dl.max_len); |
| printf("[test] sample[0]: n_tokens=%d ids[0:8]=", n0); |
| for (int i = 0; i < (n0 < 8 ? n0 : 8); i++) printf("%d ", buf[i]); |
| printf("\n"); |
| free(buf); |
|
|
| |
| int batch = 4; |
| int *indices = (int *)malloc(sizeof(int) * (size_t)dl.n_samples); |
| for (int i = 0; i < dl.n_samples; i++) indices[i] = i; |
| lal_dl_shuffle(indices, dl.n_samples, (unsigned)time(NULL)); |
|
|
| int **b_tokens = (int **)malloc(sizeof(int *) * (size_t)batch); |
| int *b_lens = (int *)malloc(sizeof(int) * (size_t)batch); |
| int br = dataloader_random_batch(&dl, indices, batch, b_tokens, b_lens); |
| printf("[test] random_batch returned %d\n", br); |
| if (br == batch) { |
| for (int i = 0; i < batch; i++) { |
| printf("[test] batch[%d] idx=%d len=%d ids[0:5]=", |
| i, indices[i], b_lens[i]); |
| int show = b_lens[i] < 5 ? b_lens[i] : 5; |
| for (int k = 0; k < show; k++) printf("%d ", b_tokens[i][k]); |
| printf("\n"); |
| } |
| } |
| dataloader_free_batch(b_tokens, batch); |
| free(b_tokens); |
| free(b_lens); |
| free(indices); |
|
|
| dataloader_free(&dl); |
| printf("[test] OK\n"); |
| return 0; |
| } |
|
|
| #endif |
|
|
| #endif |
|
|