JiRack_10b / DJL /JiRackJDLStage2TrainerMXNet .java
kgrabko's picture
Rename JiRackJDLStage2TrainerMXNet .java to DJL/JiRackJDLStage2TrainerMXNet .java
66d3698 verified
Raw
History Blame Contribute Delete
22.5 kB
/**
# =============================================================================
# COPYRIGHT © 2025-2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED.
# CMS Manhattan JiRack Technology — PATENT PENDING
#
# This code is proprietary.
# Personal and non-commercial research use is allowed.
# Any commercial use, derivative works for profit, or distribution
# requires a paid license and 5% royalty.
#
# Unauthorized commercial use is strictly prohibited.
# Contact: grabko@cmsmanhattan.com
# =============================================================================
*/
package com.cbsinc.cms.llm.ml;
import ai.djl.Device;
import ai.djl.Model;
import ai.djl.ndarray.NDArray;
import ai.djl.ndarray.NDList;
import ai.djl.ndarray.NDManager;
import ai.djl.ndarray.index.NDIndex;
import ai.djl.ndarray.types.DataType;
import ai.djl.ndarray.types.Shape;
import ai.djl.nn.Block;
import ai.djl.nn.Parameter;
import ai.djl.training.DefaultTrainingConfig;
import ai.djl.training.GradientCollector;
import ai.djl.training.Trainer;
import ai.djl.training.TrainingConfig;
import ai.djl.training.initializer.NormalInitializer;
import ai.djl.training.listener.TrainingListener;
import ai.djl.training.loss.Loss;
import ai.djl.training.optimizer.Optimizer;
import ai.djl.training.tracker.Tracker;
import ai.djl.util.Pair;
import java.io.DataInputStream;
import java.io.IOException;
import java.nio.file.Files;
import java.nio.file.Path;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.Comparator;
import java.util.List;
import java.util.Random;
import java.util.regex.Matcher;
import java.util.regex.Pattern;
import java.util.stream.Collectors;
import java.util.stream.Stream;
/**
* STAGE 2: Ternary QAT on SlimOrca — DJL port of train_slimorca_ternary_stage2.py.
* Uses the JiRack ternary JiRackJDLTernary_10b model (BitLinear + lambda) built in JiRackJDLTernary_10b .java.
*
* Same stage-2 logic:
* [S-1] Starts from the best stage-1 checkpoint (weights only — the
* optimizer starts FRESH; DJL does not serialize Adam moments, which
* matches the Python script's intent of not reusing stage-1 moments).
* [S-2] lambda = 1.0 constant from the first step. No warmup schedule.
* [S-3] Shards contain pre-built input_ids AND labels; labels are already
* masked (-100 on prompt tokens, real ids on response). The loss
* ignores -100. No reconstruction from attention masks.
* [S-4] Separate OUTPUT_DIR — never mixed with stage-1 checkpoints.
* [S-5] Fixed val set taken from SlimOrca shard 0 with a seeded
* permutation (seed 42), identical for every shard's validation.
* [S-6] Lower LR (2e-5): STE gradients at lambda=1.0 are noisier.
*
* DJL-specific adaptations (honest differences from the Python script):
* - SHARD FORMAT: DJL cannot read PyTorch .pt files. Shards must be
* NDList-encoded files named slimorca_data_<N>.ndlist, each containing
* two named INT64 arrays: "input_ids" (N, T) and "labels" (N, T),
* already padded (inputs with 0, labels with -100). A converter from
* your .pt shards is a small separate Python script (torch.load ->
* numpy -> write NDList) — say the word and I'll write it.
* - OPTIMIZER: DJL has no Adafactor. Adam with the same LR, weight decay
* 1e-4 and clip-grad 1.0 is used instead. Memory cost is higher
* (2 moments vs Adafactor's factored state).
* - PRECISION: no autocast/bf16 on the DJL CPU path; training runs FP32.
* - CHECKPOINTS: DJL saves parameters only. global_step is persisted in a
* sidecar "<name>.step" text file next to the .params file. Optimizer
* state is not persisted (fresh moments on resume).
* - GRAD ACCUMULATION: one GradientCollector per accumulation window;
* each micro-batch loss is scaled by 1/GRAD_ACCUM; trainer.step()
* applies the clipped update and zeroes gradients.
*/
public class JiRackJDLStage2TrainerMXNet {
// ========================= SETTINGS =========================
private static Path DATA_DIR = Paths.get("data/SlimOrca/shards");
private static Path OUTPUT_DIR = Paths.get("build/JRock_Ternary_SlimOrca");
/** [S-1] Best stage-1 checkpoint directory + model name (DJL .params). */
private static Path STAGE1_DIR = Paths.get("checkpoints/ternary_stage1");
private static String STAGE1_NAME = "jirack_data_22";
private static int BATCH_SIZE = 2;
private static int GRAD_ACCUM = 5;
private static float LR = 2e-5f; // [S-6]
private static final float WEIGHT_DECAY = 1e-4f;
private static final float CLIP_GRAD = 1.0f;
private static final double VAL_RATIO = 0.05;
private static final int AUTOSAVE_EVERY = 1000;
private static final long VAL_SEED = 42L;
private static final long IGNORE_INDEX = -100L;
private static final Device DEVICE = Device.cpu();
private static boolean USE_SMOKE_CONFIG = false;
// =============================================================
private static final Pattern SHARD_NUM = Pattern.compile("data_(\\d+)");
private static final Pattern CKPT_NUM = Pattern.compile("shard_(\\d+)");
private long globalStep = 0;
/* ------------------- Masked, shifted cross-entropy ------------------- */
/**
* loss = CE( logits[:, :-1, :], labels[:, 1:] ) ignoring positions where
* label == -100. Mean over non-ignored tokens. [S-3]
*/
static class MaskedShiftedCELoss extends Loss {
MaskedShiftedCELoss() {
super("MaskedShiftedCE");
}
@Override
public NDArray evaluate(NDList labels, NDList predictions) {
NDArray logits = predictions.singletonOrThrow(); // (B, T, V)
NDArray target = labels.singletonOrThrow(); // (B, T) INT64
long T = logits.getShape().get(1);
long V = logits.getShape().get(2);
// Shift: predict token t+1 from position t
NDArray shiftLogits = logits.get(new NDIndex(":, 0:" + (T - 1) + ", :"))
.reshape(-1, V); // (B*(T-1), V)
NDArray shiftLabels = target.get(new NDIndex(":, 1:" + T))
.reshape(-1); // (B*(T-1),)
NDArray mask = shiftLabels.neq(IGNORE_INDEX); // bool
NDArray maskF = mask.toType(DataType.FLOAT32, false);
// Replace -100 with 0 so gather() has a valid index; masked out later.
NDArray safeLabels = shiftLabels.mul(mask.toType(DataType.INT64, false))
.reshape(-1, 1);
NDArray logProb = shiftLogits.logSoftmax(1);
NDArray picked = logProb.gather(safeLabels, 1).reshape(-1); // (N,)
NDArray nll = picked.neg().mul(maskF);
NDArray denom = maskF.sum().maximum(1.0f);
return nll.sum().div(denom); // scalar
}
}
/* ----------------------------- Shard IO ----------------------------- */
/** Loads a shard: returns {input_ids (N,T), labels (N,T)} as INT64. */
private static NDArray[] loadShard(NDManager manager, Path shardFile) throws IOException {
try (DataInputStream dis = new DataInputStream(Files.newInputStream(shardFile))) {
NDList list = NDList.decode(manager, dis);
NDArray ids = null;
NDArray lbl = null;
for (NDArray a : list) {
if ("input_ids".equals(a.getName())) ids = a;
else if ("labels".equals(a.getName())) lbl = a;
}
if (ids == null || lbl == null) {
// Fallback: positional (0 = inputs, 1 = labels)
ids = list.get(0);
lbl = list.get(1);
}
return new NDArray[]{
ids.toType(DataType.INT64, false),
lbl.toType(DataType.INT64, false)
};
}
}
private static int shardNum(Path p) {
Matcher m = SHARD_NUM.matcher(p.getFileName().toString());
if (!m.find()) throw new IllegalArgumentException("Bad shard name: " + p);
return Integer.parseInt(m.group(1));
}
/* --------------------------- Checkpoint IO --------------------------- */
private void saveCheckpoint(Model model, Path dir, String name) throws IOException {
Files.createDirectories(dir);
model.save(dir, name);
Files.writeString(dir.resolve(name + ".step"), Long.toString(globalStep));
}
private long readStep(Path dir, String name) {
try {
return Long.parseLong(Files.readString(dir.resolve(name + ".step")).trim());
} catch (Exception e) {
return 0L;
}
}
/* ------------------------------ Training ------------------------------ */
public void train() throws Exception {
Files.createDirectories(OUTPUT_DIR);
System.out.println("Loading JiRack Ternary — Stage 2: SlimOrca...");
try (Model model = Model.newInstance("jirack_stage2", DEVICE);
NDManager dataManager = NDManager.newBaseManager(DEVICE)) {
Block block = JiRackJDLTernary_10b .buildModel();
block.setInitializer(new NormalInitializer(JiRackJDLTernary_10b .INIT_STD), Parameter.Type.WEIGHT);
model.setBlock(block);
TrainingConfig config = new DefaultTrainingConfig(new MaskedShiftedCELoss())
.optOptimizer(Optimizer.adam()
.optLearningRateTracker(Tracker.fixed(LR))
.optWeightDecays(WEIGHT_DECAY)
.optClipGrad(CLIP_GRAD)
.build())
.optDevices(new Device[]{DEVICE})
.addTrainingListeners(TrainingListener.Defaults.logging());
try (Trainer trainer = model.newTrainer(config)) {
// Initialize with a nominal shape; actual T varies per batch.
trainer.initialize(new Shape(BATCH_SIZE, 64));
JiRackJDLTernary_10b .applyDepthScaledInit(block);
// ---------- Resume logic: stage-2 ckpt > stage-1 best ----------
List<Path> stage2 = listCheckpoints(OUTPUT_DIR);
int lastDone = -1;
if (!stage2.isEmpty()) {
Path latest = stage2.get(stage2.size() - 1);
String name = stripParams(latest);
System.out.println("Resuming stage 2 from: " + name);
model.load(OUTPUT_DIR, name);
globalStep = readStep(OUTPUT_DIR, name);
Matcher m = CKPT_NUM.matcher(name);
if (m.find()) lastDone = Integer.parseInt(m.group(1));
} else {
Path s1 = STAGE1_DIR.resolve(STAGE1_NAME + "-0000.params");
if (Files.exists(s1)) {
System.out.println("Starting from stage-1 best: " + s1);
// [S-1] weights only; optimizer moments start fresh
model.load(STAGE1_DIR, STAGE1_NAME);
globalStep = 0; // stage-2 has its own step counter
} else {
System.out.println("WARNING: stage-1 checkpoint not found at "
+ s1.toAbsolutePath() + " — training from random init "
+ "(only sensible with --smoke).");
}
}
// [S-2] lambda constant 1.0 — warmup was completed in stage 1.
JiRackJDLTernary_10b .setLambda(block, 1.0f);
System.out.printf("lambda=1.0 (constant) | global_step=%d | LR=%.1e%n",
globalStep, LR);
// -------------------- Shards & fixed val set --------------------
List<Path> allShards = listShards(DATA_DIR);
if (allShards.isEmpty()) {
throw new IOException("No shards found in " + DATA_DIR.toAbsolutePath()
+ " (expected slimorca_data_<N>.ndlist)");
}
// [S-5] Fixed val set from shard 0, seeded permutation.
System.out.println("Building fixed val set from "
+ allShards.get(0).getFileName());
NDArray[] s0 = loadShard(dataManager, allShards.get(0));
long n0 = s0[0].getShape().get(0);
int valN = (int) (n0 * VAL_RATIO);
long[] perm = seededPermutation(n0, VAL_SEED);
long[] valIdx = java.util.Arrays.copyOfRange(perm, 0, valN);
long[] trainIdx0 = java.util.Arrays.copyOfRange(perm, valN, (int) n0);
NDArray valIdxArr = dataManager.create(valIdx);
NDArray valIds = s0[0].get(valIdxArr).duplicate();
NDArray valLbl = s0[1].get(valIdxArr).duplicate();
System.out.println("Fixed val set: " + valN + " examples (same for all shards)");
s0[0].close();
s0[1].close();
// -------------------------- Shard loop --------------------------
for (Path shardPath : allShards) {
int shardIdx = shardNum(shardPath);
if (shardIdx <= lastDone) {
System.out.println("Skipping already processed: "
+ shardPath.getFileName());
continue;
}
System.out.println("\nStarting shard: " + shardPath.getFileName());
try (NDManager shardManager = dataManager.newSubManager()) {
NDArray[] shard = loadShard(shardManager, shardPath);
NDArray ids = shard[0];
NDArray lbl = shard[1];
// Shard 0: exclude fixed-val rows from training.
if (shardIdx == 0) {
NDArray tIdx = shardManager.create(trainIdx0);
ids = ids.get(tIdx);
lbl = lbl.get(tIdx);
}
trainOneShard(trainer, config, block, model,
ids, lbl, shardIdx, shardManager);
// ---------------------- Validation ----------------------
System.out.println("Validating...");
float valLoss = validate(trainer, config, valIds, valLbl, shardManager);
System.out.printf(
"Shard %d — Fixed Val Loss: %.4f @ lambda=1.0000 (gstep=%d)%n",
shardIdx, valLoss, globalStep);
saveCheckpoint(model, OUTPUT_DIR,
"slimorca_ternary_shard_" + shardIdx);
System.out.println("Saved: slimorca_ternary_shard_" + shardIdx);
}
}
}
System.out.println("Stage 2 (Ternary + SlimOrca) finished!");
}
}
private void trainOneShard(Trainer trainer, TrainingConfig config, Block block,
Model model, NDArray ids, NDArray lbl,
int shardIdx, NDManager manager) throws IOException {
long n = ids.getShape().get(0);
long nBatches = n / BATCH_SIZE;
long[] order = seededPermutation(n, VAL_SEED + shardIdx + 1); // shuffle
int micro = 0;
GradientCollector gc = trainer.newGradientCollector();
boolean windowHasGrads = false;
try {
for (long b = 0; b < nBatches; b++) {
long[] rows = java.util.Arrays.copyOfRange(
order, (int) (b * BATCH_SIZE), (int) ((b + 1) * BATCH_SIZE));
NDArray rowIdx = manager.create(rows);
NDArray batchIds = ids.get(rowIdx);
NDArray batchLbl = lbl.get(rowIdx);
NDList preds = trainer.forward(new NDList(batchIds));
NDArray loss = config.getLossFunction()
.evaluate(new NDList(batchLbl), preds)
.div(GRAD_ACCUM);
float lossVal = loss.getFloat();
if (Float.isNaN(lossVal) || Float.isInfinite(lossVal)) {
// NaN/Inf: drop the whole accumulation window.
System.out.printf("%nNaN/Inf @ gstep=%d — window dropped%n", globalStep);
gc.close();
zeroGradients(block);
gc = trainer.newGradientCollector();
micro = 0;
windowHasGrads = false;
globalStep++;
continue;
}
gc.backward(loss);
windowHasGrads = true;
micro++;
if (micro == GRAD_ACCUM) {
trainer.step(); // clipped Adam update + zero grads
gc.close();
gc = trainer.newGradientCollector();
micro = 0;
windowHasGrads = false;
if (AUTOSAVE_EVERY > 0 && globalStep > 0
&& globalStep % AUTOSAVE_EVERY < GRAD_ACCUM) {
saveCheckpoint(model, OUTPUT_DIR, "autosave_latest");
System.out.printf("autosave @ gstep=%d%n", globalStep);
}
}
if (b % 10 == 0) {
System.out.printf("Shard %d [%d/%d] loss=%.4f lambda=1.0000 gstep=%d%n",
shardIdx, b, nBatches, lossVal * GRAD_ACCUM, globalStep);
}
globalStep++;
}
} finally {
// Leftover partial window: apply it (small final step) or discard.
if (windowHasGrads) {
trainer.step();
}
gc.close();
}
}
private float validate(Trainer trainer, TrainingConfig config,
NDArray valIds, NDArray valLbl, NDManager manager) {
long n = valIds.getShape().get(0);
long nBatches = Math.max(1, n / BATCH_SIZE);
double total = 0.0;
int steps = 0;
for (long b = 0; b < nBatches; b++) {
NDArray batchIds = valIds.get(new NDIndex(
(b * BATCH_SIZE) + ":" + Math.min(n, (b + 1) * BATCH_SIZE) + ", :"));
NDArray batchLbl = valLbl.get(new NDIndex(
(b * BATCH_SIZE) + ":" + Math.min(n, (b + 1) * BATCH_SIZE) + ", :"));
NDList preds = trainer.evaluate(new NDList(batchIds));
float v = config.getLossFunction()
.evaluate(new NDList(batchLbl), preds).getFloat();
if (Float.isFinite(v)) {
total += v;
steps++;
}
}
return steps > 0 ? (float) (total / steps) : Float.POSITIVE_INFINITY;
}
/* ------------------------------ Helpers ------------------------------ */
private static void zeroGradients(Block block) {
for (Pair<String, Parameter> p : block.getParameters()) {
NDArray arr = p.getValue().getArray();
if (arr.hasGradient()) {
NDArray g = arr.getGradient();
g.muli(0);
}
}
}
private static long[] seededPermutation(long n, long seed) {
long[] idx = new long[(int) n];
for (int i = 0; i < n; i++) idx[i] = i;
Random rnd = new Random(seed);
for (int i = (int) n - 1; i > 0; i--) {
int j = rnd.nextInt(i + 1);
long t = idx[i];
idx[i] = idx[j];
idx[j] = t;
}
return idx;
}
private static List<Path> listShards(Path dir) throws IOException {
if (!Files.isDirectory(dir)) return List.of();
try (Stream<Path> s = Files.list(dir)) {
return s.filter(p -> p.getFileName().toString().matches("slimorca_data_\\d+\\.ndlist"))
.sorted(Comparator.comparingInt(GPTStage2Trainer::shardNum))
.collect(Collectors.toList());
}
}
private static List<Path> listCheckpoints(Path dir) throws IOException {
if (!Files.isDirectory(dir)) return List.of();
try (Stream<Path> s = Files.list(dir)) {
return s.filter(p -> p.getFileName().toString()
.matches("slimorca_ternary_shard_\\d+-\\d+\\.params"))
.sorted(Comparator.comparingInt(p -> {
Matcher m = CKPT_NUM.matcher(p.getFileName().toString());
m.find();
return Integer.parseInt(m.group(1));
}))
.collect(Collectors.toList());
}
}
private static String stripParams(Path p) {
String f = p.getFileName().toString();
return f.substring(0, f.lastIndexOf('-')); // drop "-0000.params"
}
/* -------------------------------- Main -------------------------------- */
public static void main(String[] args) {
System.out.println("java -cp Jirackkit.jar com.cbsinc.cms.llm.ml.GPTStage2Trainer "
+ "[batch] [gradAccum] [lr] [dataDir] [outDir] [stage1Dir] [stage1Name] [--smoke]");
if (args.length > 0) BATCH_SIZE = Integer.parseInt(args[0]);
if (args.length > 1) GRAD_ACCUM = Integer.parseInt(args[1]);
if (args.length > 2) LR = Float.parseFloat(args[2]);
if (args.length > 3) DATA_DIR = Paths.get(args[3]);
if (args.length > 4) OUTPUT_DIR = Paths.get(args[4]);
if (args.length > 5) STAGE1_DIR = Paths.get(args[5]);
if (args.length > 6) STAGE1_NAME = args[6];
for (String a : args) {
if ("--smoke".equals(a)) USE_SMOKE_CONFIG = true;
}
if (USE_SMOKE_CONFIG) {
System.out.println("SMOKE MODE: tiny model config");
JiRackJDLTernary_10b .smokeConfig();
}
try {
new GPTStage2Trainer().train();
} catch (Exception e) {
System.err.println("Stage-2 training failed: " + e.getMessage());
e.printStackTrace();
}
}
}