-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathprepro.js
More file actions
171 lines (142 loc) · 5.93 KB
/
Copy pathprepro.js
File metadata and controls
171 lines (142 loc) · 5.93 KB
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
// Preprocessing: Download MNIST, resize to 16x16, save binary
// Faithful reproduction of Karpathy's prepro.py
// Dataset sourced from HuggingFace (ylecun/mnist) via IDX format mirror
import fs from 'fs';
import path from 'path';
import { gunzipSync } from 'zlib';
import { RNG } from './lib/rng.js';
const DATA_DIR = './data';
// MNIST IDX files hosted on Google Cloud (same data as HuggingFace ylecun/mnist)
const URLS = {
trainImages: 'https://storage.googleapis.com/cvdf-datasets/mnist/train-images-idx3-ubyte.gz',
trainLabels: 'https://storage.googleapis.com/cvdf-datasets/mnist/train-labels-idx1-ubyte.gz',
testImages: 'https://storage.googleapis.com/cvdf-datasets/mnist/t10k-images-idx3-ubyte.gz',
testLabels: 'https://storage.googleapis.com/cvdf-datasets/mnist/t10k-labels-idx1-ubyte.gz',
};
const N_TRAIN = 7291;
const N_TEST = 2007;
// Bilinear interpolation resize (align_corners=false, matching PyTorch F.interpolate)
function resizeBilinear(src, srcH, srcW, dstH, dstW) {
const dst = new Float64Array(dstH * dstW);
for (let dy = 0; dy < dstH; dy++) {
for (let dx = 0; dx < dstW; dx++) {
const sx = (dx + 0.5) * (srcW / dstW) - 0.5;
const sy = (dy + 0.5) * (srcH / dstH) - 0.5;
const x0 = Math.floor(sx), y0 = Math.floor(sy);
const xf = sx - x0, yf = sy - y0;
const clamp = (v, lo, hi) => Math.max(lo, Math.min(hi, v));
const cx0 = clamp(x0, 0, srcW - 1), cx1 = clamp(x0 + 1, 0, srcW - 1);
const cy0 = clamp(y0, 0, srcH - 1), cy1 = clamp(y0 + 1, 0, srcH - 1);
const v00 = src[cy0 * srcW + cx0];
const v01 = src[cy0 * srcW + cx1];
const v10 = src[cy1 * srcW + cx0];
const v11 = src[cy1 * srcW + cx1];
dst[dy * dstW + dx] = (1 - yf) * ((1 - xf) * v00 + xf * v01)
+ yf * ((1 - xf) * v10 + xf * v11);
}
}
return dst;
}
// Download and decompress a .gz file, return Buffer
async function downloadGz(url, label) {
const cachePath = path.join(DATA_DIR, path.basename(url));
if (fs.existsSync(cachePath)) {
console.log(` [cached] ${label}`);
return gunzipSync(fs.readFileSync(cachePath));
}
console.log(` Downloading ${label}...`);
const resp = await fetch(url);
if (!resp.ok) throw new Error(`Failed to download: ${resp.status}`);
const gz = Buffer.from(await resp.arrayBuffer());
fs.writeFileSync(cachePath, gz);
console.log(` ${(gz.byteLength / 1024 / 1024).toFixed(1)} MB`);
return gunzipSync(gz);
}
// Parse MNIST IDX image file -> { count, rows, cols, pixels: Uint8Array }
function parseImages(buf) {
const magic = buf.readUInt32BE(0);
if (magic !== 2051) throw new Error(`Bad image magic: ${magic}`);
const count = buf.readUInt32BE(4);
const rows = buf.readUInt32BE(8);
const cols = buf.readUInt32BE(12);
const pixels = new Uint8Array(buf.buffer, buf.byteOffset + 16, count * rows * cols);
return { count, rows, cols, pixels };
}
// Parse MNIST IDX label file -> { count, labels: Uint8Array }
function parseLabels(buf) {
const magic = buf.readUInt32BE(0);
if (magic !== 2049) throw new Error(`Bad label magic: ${magic}`);
const count = buf.readUInt32BE(4);
const labels = new Uint8Array(buf.buffer, buf.byteOffset + 8, count);
return { count, labels };
}
// Save preprocessed data to binary file
function saveData(filepath, images, labels) {
const n = images.length;
const sampleSize = 256 + 10; // 16*16 pixels + 10 label values
const headerSize = 8;
const buf = Buffer.alloc(headerSize + n * sampleSize * 8);
buf.writeUInt32LE(n, 0);
buf.writeUInt32LE(0, 4); // reserved
const floats = new Float64Array(buf.buffer, buf.byteOffset + headerSize, n * sampleSize);
for (let i = 0; i < n; i++) {
const offset = i * sampleSize;
for (let j = 0; j < 256; j++) {
floats[offset + j] = images[i][j];
}
for (let j = 0; j < 10; j++) {
floats[offset + 256 + j] = labels[i][j];
}
}
fs.writeFileSync(filepath, buf);
console.log(` Saved ${n} samples to ${filepath} (${(buf.byteLength / 1024 / 1024).toFixed(1)} MB)`);
}
async function main() {
fs.mkdirSync(DATA_DIR, { recursive: true });
console.log('Downloading MNIST from Google Cloud (same as HuggingFace ylecun/mnist)...');
const [trainImgBuf, trainLblBuf, testImgBuf, testLblBuf] = await Promise.all([
downloadGz(URLS.trainImages, 'train images'),
downloadGz(URLS.trainLabels, 'train labels'),
downloadGz(URLS.testImages, 'test images'),
downloadGz(URLS.testLabels, 'test labels'),
]);
const trainImg = parseImages(trainImgBuf);
const trainLbl = parseLabels(trainLblBuf);
const testImg = parseImages(testImgBuf);
const testLbl = parseLabels(testLblBuf);
console.log(`\nParsed: ${trainImg.count} train, ${testImg.count} test (${trainImg.rows}x${trainImg.cols})`);
const rng = new RNG(1337);
// Process each split
for (const [name, img, lbl, nSelect] of [
['train', trainImg, trainLbl, N_TRAIN],
['test', testImg, testLbl, N_TEST],
]) {
console.log(`\nProcessing ${name}: selecting ${nSelect} from ${img.count}...`);
const perm = rng.permutation(img.count);
const selected = perm.slice(0, nSelect);
const images = [];
const labels = [];
const pixelSize = img.rows * img.cols;
for (const idx of selected) {
// Extract 28x28 image, normalize to [-1, +1]
const src = new Float64Array(pixelSize);
const pixelOffset = idx * pixelSize;
for (let i = 0; i < pixelSize; i++) {
src[i] = img.pixels[pixelOffset + i] / 127.5 - 1.0;
}
// Resize 28x28 -> 16x16
const resized = resizeBilinear(src, 28, 28, 16, 16);
images.push(resized);
// Encode label: 10-dim vector, all -1 except +1 at correct index
const labelVec = new Float64Array(10).fill(-1.0);
labelVec[lbl.labels[idx]] = 1.0;
labels.push(labelVec);
}
saveData(path.join(DATA_DIR, `${name}1989.bin`), images, labels);
}
console.log('\nPreprocessing complete!');
}
main().catch(err => {
console.error('Error:', err);
process.exit(1);
});