All posts

PMPP Chapter 20 Notes

PMPP Chapter 20 Notes

Note

This is not a complete notes for chapters, but mostly prompted by working on things inside the chapter

History

  • Google birthed transformers in 2017

    • Its application on BERT led to unprecendented natural language generation results
  • New backbone of deep learning

  • LLM vs Multimodel models

  • Autoregression already existed (i.e. LSTM, etc)

  • GPT 1

    • Decoder only transformers on wikipedia text
  • GPT 2

    • If you train on the right data, you get implicit multitask learning - i.e. the model understands multiple types of tasks, you just need the right adapters for them
  • GPT 3

    • The task can just be provided or described in the same token stream and the model will handle the rest

Terminology

alt text

  • Context (for chatgbots, etc)

    • Typically composed of system prompt, user input, toens already generated by the model
  • Feed forward network

  • Architecture

    • Typically when ML practitioners refer to architecture, they mean coarsely to i.e.
      • decoder only transformer
      • 32 layers
      • residual with 4096
      • etc
  • Parameterization

    • i.e. is weight stored directly as W or UV?
    • is positive scalr stored directly or as eρe^\rho
    • is weight vector stored directly or as magnitude times normalized direction?
    • Are input/output embeddings separate or tied?
    • is a residual branch represented as F(x) or alpha F(x) with learned alpha?
    • Are attention projections independent, shared, or low rank?
  • Encoding

    • extremely overloaded term
    • Meaning 1:
      • Encoding text
      • string -> token IDs
    • Meaning 2:
      • Transformer encoder
      • refers to the "neural net encoder" part of the original transformer paper
      • Mapped ENTIRE SEQUENCE to contextual representations Z
        • Decoder then generated target sequence one symbol at a time using both
          • target tokens generated so far and
          • encoder's output Z
  • Embedding

    • token IDs -> embedding vectors
  • Positional encoding

    • position -> numerical positional information
  • Contextual encoding

    • hiddent vectors come to represent the surrounding prefix
  • Transformer encoder

    • a particular noncausal source proceeing architecture
  • Decoding

    • unfortunately, also overloaded term. Can mean
    • choosing tokens from model probabilities
    • converting token ids back to text
    • or operating target side network of an encoder decoder architecture

Numbers

  • LLM
    • Billions to trillions of parameters

Example

  • Why just K and V? why not Q?
    • Let's walk through what happens when N→N+1N \to N+1
    • alt text

Architecture

  • Attention

    • To some extent resembles a classificaiton process within a neurla net
      • Generate probabilities with softmax
      • Indication of relative importance
  • Types of tasks

    • Transformers that perform discriminative rather than generative tasks do not need a decoder
      • They do not need to generate anything and auto regress on itself
    • Decod
  • Things that affect model architecture

    • Hard information constraints
      • i.e. causal mask
        • h_t can only depend on 1,...,t
    • How much memory/compute is needed to learn something
    • Worth noting that two parameterizations with exact same represented function class can train differently:
      • i.e. W=UVW = UV

      • ΔW=−η∇WL\Delta W = -\eta \nabla_W L

      • But consider what happens to Δ(UV)\Delta(UV) in the factorized case

        • ΔU=−η(∇UVL)VT\Delta U = -\eta (\nabla_{UV} L) V^T and ΔV=−ηUT(∇UVL)\Delta V = -\eta U^T(\nabla_{UV}L)
        • This means in the factorized case, Δ(UV)=~−η((∇UVL)VTV+UUT(∇UVL))\Delta(UV) \tilde{=} -\eta((\nabla_{UV}L)V^TV + UU^T(\nabla_{UV}L))
        • Note the above does NOT equal −η∇WL-\eta \nabla_W L
    • Also the initial distributions, optimizers, normalization, regularization, parameterization can all affect resultant fits
      • On a first order, useful architectures make the desired regularities relatively cheap and undesirable memorization or patterns expensive
  • for X(l)∈RB×N×DX^{(l)} \in \RR^{B \times N \times D}

    • Note Input and Output of each block has the same shape, but within each block
    • There is cross-N contamination (causal) in N
    • MLP mixes across D
    • Residual stream carries exact same info across
  • Note that WQ,WK,WVW_Q, W_K, W_V have the same shape. The difference in their function is only determined by the differences in weight

  • Unembedding

    • Turning embedding back into text

Rough architecture of a GPU

                    HBM stacks
                        │
              HBM PHYs/controllers
                        │
              physically sliced L2
                        │
              on-chip fabric / crossbar
                        │
              GPU Processing Clusters
                        │
          ┌─────────────┴─────────────┐
          │                           │
         SM                          SM
  ┌────────────────┐          ┌────────────────┐
  │ warp schedulers│          │ warp schedulers│
  │ register files │          │ register files │
  │ Tensor Cores   │          │ Tensor Cores   │
  │ FP/INT ALUs    │          │ FP/INT ALUs    │
  │ LSU / SFU      │          │ LSU / SFU      │
  │ L1/shared SRAM │          │ L1/shared SRAM │
  │ barriers/TMA   │          │ barriers/TMA   │
  └────────────────┘          └────────────────┘

Rough architecture of a TPU

                         HBM
                          │
                    asynchronous DMA
                          │
              ┌───────────┴───────────┐
              │                       │
             VMEM                    SMEM
        vector SRAM              scalar SRAM
              │                       │
             VREG                    SREG
              │                       │
       ┌──────┼──────────┐            │
       │      │          │            │
      MXU    VPU        XLU       scalar unit
   matmul   vector    transpose/   control,
                    permutation    addresses

Semantic interpretation

  • Aij(h)A_{ij}^{(h)}

    • For head h, how important is token at source position jj when updating ii?
  • Q∈RN×dQ \in \RR^{N \times d}

  • K∈RN×dK \in \RR^{N \times d}

  • V∈RN×dV \in \RR^{N \times d}

  • X∈RN×dX \in \RR^{N \times d}

  • WQ∈Rd×dW_Q \in \RR^{d \times d}

  • WK∈Rd×dW_K \in \RR^{d \times d}

  • WV∈Rd×dW_V \in \RR^{d \times d}

  • Softmax

  • Each QQ's row is a token position.

  • Each KTK^T's column is a token position

  • Dotting QQ's ii'th row with KK's jj'th column gives

    • Sij=How relevant is j for determining iS_{ij} = \text{How relevant is j for determining i}
    • Can see it's i for j because j≤ij \le i in the mask
    • alt text
    • Softwax over each row -> (semantically probability) coefficients for position jj for ii
      • The element at (i,j)(i,j) in (QKTd+M)(\frac{QK^T}{\sqrt{d}} + M) is semantically seen as the logit lr,cl_{r,c}
      • mr,cm_{r,c} is subtracted from lr,cl_{r,c} to prevent overflow. Both numerator and denominator so it cancels out
      • i.e.
        • Pr,c=elr,c−mr∑j=1Nelr,j−mrP_{r,c} = \frac{e^{l_{r,c} - m_{r}}}{\sum_{j=1}^N e^{l_{r,j} - m_r}}
    • Multiply by VV
      • Interpreted as each row in SS:
        • for row ii,
          • for column jj
            • Add SijS_{ij} times VV's jj'th row to ii
  • P=softmax⁡(QKT+M)P = \operatorname{softmax}(QK^T + M)

  • OO is the output

  • Note in reality there is a sacling factor of 1d\frac{1}{\sqrt{d}}

Numbers

  • Indices

    • i,j∈{1,...,n}i,j \in \{1, ..., n\} (TOKEN POSITION)
    • r,s∈{1,...,d}r, s \in \{1, ..., d\} (HIDDEN DIMENSION POSITION)
    • h∈{1,...,h}h \in \{1, ..., h\} (HEAD, AXIS)
    • Attention matrix - one for each head
    • Rn×n\RR^{n \times n}
  • Dimensions

    • X∈Rn×d\mathbf{X} \in \RR^{n \times d}

      X=[x⃗1x⃗2⋮x⃗n]∈Rn×d.\mathbf{X} = \begin{bmatrix} \vec{x}_1 \\ \vec{x}_2 \\ \vdots \\ \vec{x}_n \end{bmatrix} \in \RR^{n \times d}.
    • A(h)∈Rn×n\mathbf{A}^{(h)} \in \RR^{n \times n}

  • Note, if batches exist, typically do

    • X(l)∈RB×N×DX^{(l)} \in \RR^{B \times N \times D}
      • l - depth l circuit
      • B - independent sequences in the batch
      • N - token positions in each sequence
      • D - dmodeld_{\text{model}} - feature channels carried by each token
    • Block interfaces keep the same shape
    • X←X+Attention⁡(Norm⁡(X))X \leftarrow X + \operatorname{Attention}(\operatorname{Norm}(X))