S-exp GPU
All pages
Docs · The run fileMarkdown

Generation inside the step

defgenerate makes records on the device, inside each training step, before the step's training graph reads them: a student sampling its own rows, a frozen teacher labelling them, a pool scored and cut down to the rows worth a gradient. It is a loop of ordinary graphs the runtime runs; no op carries a subgraph, and nothing is differentiated through it.

(defgenerate :trips 6
             :init (lambda (batch ctx) (list :row ... :stride ...))
             :step (lambda (state trip ctx) (list :row ... :stride ...))
             :finish (lambda (state ctx) (list :sample ... :label ...)))
keywordmeaning
:tripshow many times :step runs, 1 or more
:initthe batch, as prepare receives it, and the step's context, to the state: a plist of tensors, a keyword before each
:stepthe state, the trip as an i64 scalar and the context, to the next state, with the keys, shapes and dtypes :init made
:finishoptional: the last state and the context, once, to any plist of tensors

Each training microbatch runs :init on its batch, :step :trips times and :finish once, the state staying on the device. What :finish returned, or the last state without one, joins the batch: prepare reads each entry with (field batch :name) as it reads a loader field, so a generated name may not be a field of the training loader or of any pass's loader. An evaluation pass whose loader has the training loader's fields and batch size runs the generator on its batches too; a pass whose loader differs does not, and its prepare reading a generated name is E-CONTRACT-005.

The state has static shapes. A sequence written one position per trip is a fixed-size tensor written with where at the trip's column; a cache is the same. The fixture crates/cli/tests/fixtures/tiny-generate.sx has the student write each row greedily, one position per trip, and :finish label every position by the rule it is learning:

:step (lambda (state trip ctx)
        (let* ((row (getf state :row))
               (at (+ (cast trip :i32) 2))
               (column (iota [1 seq] :axis 1))
               (guess (sum (where (= column (- at 1)) (argmax (model row)) 0)
                           :axes [1]
                           :keepdims true)))
          (list :row (where (= column at) guess row)
                :stride (getf state :stride))))

What a generator may read

  • The batch's fields in :init, the state in :step and :finish, and every parameter, frozen or not: the model model binds, or any model reachable from it, runs inside a generator as it does in objective, at the parameters the step started from.
  • The context: the builtin keys of ctx only, and the trip as :step's second argument. A counter is not among them: prepare declares counters, and it is traced after the generator, so (getf ctx :tokens) is nil there. A generator that reads :microbatch, :records or :stage-records keeps the run unstacked, since those move between the microbatches one stacked call holds.
  • A draw is a pure function of the seed, its place and its salt; see random draws. An unsalted draw is the same at every step and microbatch. Salted with (+ (* (getf ctx :step) 65536) (getf ctx :microbatch)) it is new for every microbatch, and the :microbatch keeps the run unstacked; a salt of :step alone is new every step and stacks.

What a generator may not do

  • Report: a metric inside it is E-GEN-004; return the value in the state and report it from objective. A diagnostic inside it, which a model it runs may well call, is never selected and is dropped, as an unselected diagnostic is anywhere.
  • Declare a counter: counter belongs to prepare (E-CONTRACT-010).
  • Change the state's keys, shapes or dtypes from trip to trip (E-GEN-002), or return anything but a plist of tensors (E-GEN-001). Under :bf16 a state entry keeps the dtype it was made with, so a cache can be bf16 from trip to trip.

A phase of the step

  • Stacking. A run that stacks microbatches stacks the generator with the same degree: each copy has its own batch and state, the parameters and the step's context scalars are shared, and the bits are those of the microbatches one at a time. The cost model prices a step's generation and training calls together; see devices.
  • Memory. Each generator graph is a phase of the memory plan: its live set beside the batch :init read and the state the trips carry, and what it made lives until the training call that reads it. The report's memory line shows the largest, generation 0.4 GiB.
  • Signals. The loop looks for a signal before every trip; a stop abandons the step in flight, and the resume takes it again. See checkpoints.
  • Reports. sexpgpu explain lists generate.init, generate.step and generate.finish with the graphs and the trips under the training step; the lowering report has a generate row, generate stacked when the run stacks, and a generation line.
  • Determinism. A generator is part of the file, so it adds nothing to a checkpoint's identity, and the trip is an input like the step; see determinism.

Related: the run file, diagnostic codes.