/** # Copyright (c) 2025-2026 CMS Manhattan # All rights reserved. # Author: Konstantin Vladimirovich Grabko # Email: grabko@cmsmanhattan.com # Phone: +1(516)777-0945 */ package com.cbsinc.cms.llm.ml; import ai.djl.Device; import ai.djl.Model; import ai.djl.engine.Engine; 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.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 — Java on the PYTORCH ENGINE (libtorch). * * This is the "PyTorch on Java" build: DJL's Java API executing every tensor * op on native libtorch (the same C++ kernels CPython PyTorch uses). Ops this * model needs that only the PyTorch engine provides reliably: stopGradient * (the STE in BitLinear), stepped slicing (interleaved RoPE 0::2 / 1::2), * gather (masked cross-entropy), stack, and CUDA execution. * * --------------------------- DEPENDENCIES (Gradle) --------------------------- * implementation platform("ai.djl:bom:0.36.0") * implementation "ai.djl:api" * implementation "ai.djl.pytorch:pytorch-engine" * // Pick ONE native runtime: * runtimeOnly "ai.djl.pytorch:pytorch-native-cu124::linux-x86_64" // NVIDIA GPU * // runtimeOnly "ai.djl.pytorch:pytorch-native-cpu::linux-x86_64" // CPU only * runtimeOnly "ai.djl.pytorch:pytorch-jni" * * Maven uses the same artifact IDs. If no native runtime is bundled, the * first run downloads libtorch automatically. * * ------------------------------- JVM FLAGS ------------------------------- * -Dai.djl.default_engine=PyTorch * -Xmx8g (JVM heap holds only Java objects; tensors live in NATIVE * memory, so a huge -Xmx is neither needed nor helpful) * * ------------------------------ SHARD FORMAT ------------------------------ * DJL cannot unpickle .pt files. Convert each PyTorch shard once with this * small script (run in the Python env that created the shards): * * # pt_to_ndlist.py * import torch, numpy as np, sys * d = torch.load(sys.argv[1], map_location="cpu", weights_only=False) * ids = torch.nn.utils.rnn.pad_sequence(d["input_ids"], batch_first=True, * padding_value=0).to(torch.int64) * lbl = torch.nn.utils.rnn.pad_sequence(d["labels"], batch_first=True, * padding_value=-100).to(torch.int64) * np.save(sys.argv[1] + ".ids.npy", ids.numpy()) * np.save(sys.argv[1] + ".lbl.npy", lbl.numpy()) * * The trainer reads either slimorca_data_.ndlist (named NDList with * "input_ids"/"labels") or the .npy pair next to slimorca_data_.pt. * ============================================================================= */ public class JiRackJDLStage2TrainerPt { // ========================= SETTINGS ========================= private static Path DATA_DIR = Paths.get("data/SlimOrca/shards"); private static Path OUTPUT_DIR = Paths.get("build/JRock_Ternary_SlimOrca_pt"); 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 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; private Device device; /* --------------------- Engine & device selection --------------------- */ private Device selectDevice() { Engine engine = Engine.getInstance(); if (!"PyTorch".equals(engine.getEngineName())) { throw new IllegalStateException( "This trainer requires the PyTorch engine (found: " + engine.getEngineName() + "). Add ai.djl.pytorch:pytorch-engine " + "to the classpath and set -Dai.djl.default_engine=PyTorch."); } System.out.println("Engine: PyTorch " + engine.getVersion()); int gpus = engine.getGpuCount(); if (gpus > 0) { System.out.println("CUDA devices: " + gpus + " -> using gpu(0)"); return Device.gpu(0); } System.out.println("No CUDA device -> CPU (use --smoke for the tiny config)"); return Device.cpu(); } /* ------------------- Masked, shifted cross-entropy ------------------- */ /** CE(logits[:, :-1, :], labels[:, 1:]) ignoring label == -100. [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); NDArray shiftLogits = logits.get(new NDIndex(":, 0:" + (T - 1) + ", :")) .reshape(-1, V); NDArray shiftLabels = target.get(new NDIndex(":, 1:" + T)) .reshape(-1); NDArray mask = shiftLabels.neq(IGNORE_INDEX); NDArray maskF = mask.toType(DataType.FLOAT32, false); // Replace -100 with 0 so gather() has a valid index; masked later. NDArray safeLabels = shiftLabels.mul(mask.toType(DataType.INT64, false)) .reshape(-1, 1); // Loss math in FP32 for stability regardless of model dtype. NDArray logProb = shiftLogits.toType(DataType.FLOAT32, false).logSoftmax(1); NDArray picked = logProb.gather(safeLabels, 1).reshape(-1); NDArray nll = picked.neg().mul(maskF); return nll.sum().div(maskF.sum().maximum(1.0f)); } } /* ----------------------------- Shard IO ----------------------------- */ private static NDArray[] loadShard(NDManager manager, Path shardFile) throws IOException { String fn = shardFile.getFileName().toString(); if (fn.endsWith(".ndlist")) { 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) { ids = list.get(0); lbl = list.get(1); } return new NDArray[]{ids.toType(DataType.INT64, false), lbl.toType(DataType.INT64, false)}; } } // .npy pair produced by the converter in the header Path idsNpy = shardFile.resolveSibling(fn + ".ids.npy"); Path lblNpy = shardFile.resolveSibling(fn + ".lbl.npy"); if (Files.exists(idsNpy) && Files.exists(lblNpy)) { NDArray ids = manager.decode(Files.readAllBytes(idsNpy)); NDArray lbl = manager.decode(Files.readAllBytes(lblNpy)); return new NDArray[]{ids.toType(DataType.INT64, false), lbl.toType(DataType.INT64, false)}; } throw new IOException("Shard not readable: " + shardFile + " (need .ndlist or the .npy pair — see the converter in the header)"); } 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 { device = selectDevice(); Files.createDirectories(OUTPUT_DIR); System.out.println("Loading JiRack Ternary — Stage 2 (PyTorch engine)..."); try (Model model = Model.newInstance("jirack_stage2", device); NDManager rootManager = 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)) { trainer.initialize(new Shape(BATCH_SIZE, 64)); JiRackJDLTernary_10b.applyDepthScaledInit(block); // ---------- Resume: stage-2 checkpoint > stage-1 best ---------- List 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); model.load(STAGE1_DIR, STAGE1_NAME); // weights only [S-1] globalStep = 0; // stage-2 has its own step counter } else { System.out.println("WARNING: no stage-1 checkpoint at " + s1.toAbsolutePath() + " — random init (only sensible with --smoke)."); } } // [S-2] lambda constant 1.0 — warmup was done in stage 1. JiRackJDLTernary_10b.setLambda(block, 1.0f); System.out.printf( "lambda=1.0 (constant) | global_step=%d | LR=%.1e | device=%s%n", globalStep, LR, device); // -------------------- Shards & fixed val set -------------------- List allShards = listShards(DATA_DIR); if (allShards.isEmpty()) { throw new IOException("No shards in " + DATA_DIR.toAbsolutePath()); } // [S-5] Fixed val set from shard 0, seed 42. System.out.println("Building fixed val set from " + allShards.get(0).getFileName()); long[] trainIdx0; NDArray valIds; NDArray valLbl; { NDArray[] s0 = loadShard(rootManager, 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); trainIdx0 = java.util.Arrays.copyOfRange(perm, valN, (int) n0); NDArray vi = rootManager.create(valIdx); valIds = s0[0].get(vi).duplicate(); valLbl = s0[1].get(vi).duplicate(); s0[0].close(); s0[1].close(); vi.close(); System.out.println("Fixed val set: " + valN + " examples (same for all shards)"); } // -------------------------- 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()); // Shard-scoped native memory: freed when this manager closes // (the GPU equivalent of del payload + empty_cache + gc). try (NDManager shardManager = rootManager.newSubManager()) { NDArray[] shard = loadShard(shardManager, shardPath); NDArray ids = shard[0]; NDArray lbl = shard[1]; 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); 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, PyTorch engine) finished!"); } } private void trainOneShard(Trainer trainer, TrainingConfig config, Block block, Model model, NDArray ids, NDArray lbl, int shardIdx, NDManager shardManager) 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++) { float lossVal; boolean dropped = false; // Per-batch scope: every intermediate tensor of this step is // freed when the sub-manager closes. On GPU this is what // prevents OOM creep that a shard-level manager can't stop. try (NDManager batchManager = shardManager.newSubManager()) { long[] rows = java.util.Arrays.copyOfRange( order, (int) (b * BATCH_SIZE), (int) ((b + 1) * BATCH_SIZE)); NDArray rowIdx = batchManager.create(rows); NDArray batchIds = ids.get(rowIdx); NDArray batchLbl = lbl.get(rowIdx); batchIds.attach(batchManager); batchLbl.attach(batchManager); NDList preds = trainer.forward(new NDList(batchIds)); NDArray loss = config.getLossFunction() .evaluate(new NDList(batchLbl), preds) .div(GRAD_ACCUM); lossVal = loss.getFloat(); if (Float.isNaN(lossVal) || Float.isInfinite(lossVal)) { dropped = true; } else { gc.backward(loss); windowHasGrads = true; micro++; } } if (dropped) { // 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; } if (micro == GRAD_ACCUM) { trainer.step(); // clipped Adam update + zeroed 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 { if (windowHasGrads) { trainer.step(); // apply the leftover partial window } gc.close(); } } private float validate(Trainer trainer, TrainingConfig config, NDArray valIds, NDArray valLbl, NDManager shardManager) { 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++) { try (NDManager batchManager = shardManager.newSubManager()) { long lo = b * BATCH_SIZE; long hi = Math.min(n, (b + 1) * BATCH_SIZE); NDArray batchIds = valIds.get(new NDIndex(lo + ":" + hi + ", :")); NDArray batchLbl = valLbl.get(new NDIndex(lo + ":" + hi + ", :")); batchIds.attach(batchManager); batchLbl.attach(batchManager); 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 p : block.getParameters()) { NDArray arr = p.getValue().getArray(); if (arr.hasGradient()) { arr.getGradient().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 listShards(Path dir) throws IOException { if (!Files.isDirectory(dir)) return List.of(); try (Stream s = Files.list(dir)) { return s.filter(p -> { String f = p.getFileName().toString(); return f.matches("slimorca_data_\\d+\\.ndlist") || f.matches("slimorca_data_\\d+\\.pt"); }) .sorted(Comparator.comparingInt(GPTStage2TrainerPt::shardNum)) .collect(Collectors.toList()); } } private static List listCheckpoints(Path dir) throws IOException { if (!Files.isDirectory(dir)) return List.of(); try (Stream 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 -Dai.djl.default_engine=PyTorch -cp Jirackkit.jar " + "com.cbsinc.cms.llm.ml.GPTStage2TrainerPt " + "[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 GPTStage2TrainerPt().train(); } catch (Exception e) { System.err.println("Stage-2 training failed: " + e.getMessage()); e.printStackTrace(); } } }