I read the paper and code so you don't have to: Self-Adaptive Large Language Models

| January 18, 2025

How SakanaAI's Transformer² adapts LLM weights at inference time via SVD fine-tuning, RL, and expert weight combination, mapped from paper to code.

This post contains affiliate links for tools I use in production. If you buy through them I earn a commission at no extra cost to you. Recommendations are based on my own experience.

Table of Contents


This post is based on the research and code from SakanaAI’s Transformer²: Self-Adaptive LLMs project. The paper describes a method for adapting a base LLM’s weights to new tasks without full retraining, and the accompanying codebase implements the method described below.

Paper: Transformer²: Self-Adaptive LLMs

Code: SakanaAI’s Transformer² GitHub Repository

Blog: SakanaAI’s Transformer² Blog Post


1. The Mathematical Backbone: SVF, RL, and Weighted Combinations

Singular Value Fine-tuning (SVF)

SVF leverages Singular Value Decomposition (SVD) to fine-tune LLMs efficiently:

  • SVD of a Weight Matrix:

    W=U⋅Σ⋅VTW = U \cdot \Sigma \cdot V^T

    where Σ=diag(σ1,…,σr)\Sigma = \text{diag}(\sigma_1, \dots, \sigma_r) contains singular values σi\sigma_i.

  • Fine-tuning with a Mask: Adjust the singular values with a learnable vector z\mathbf{z} (or mask m\mathbf{m}) to obtain new weights:

    W′=U⋅diag(σi×mi)⋅VTW' = U \cdot \text{diag}(\sigma_i \times m_i) \cdot V^T

    The mask m\mathbf{m} scales each singular value, modulating the contribution of each singular component in W′W'.

Python-to-Math Variable Correspondence:

  • W in the code corresponds to a weight matrix of the model.
  • U, S, and V (extracted as decomposed_params[f"{param_name}.U"], ...S, ...V in the code) represent the components from the SVD of WW.
  • mm (from policy.get_mask(...)) is the mask m\mathbf{m}.
  • The resulting recomposed matrix W' is stored in variables like new_params[k].

Code Implementation:

  • compose_new_params Function (utils.py):
    mm = policy.get_mask(learnable_params[param_name])
    return (
    decomposed_params[f"{param_name}.U"]
    @ torch.diag_embed(decomposed_params[f"{param_name}.S"] * mm)
    @ decomposed_params[f"{param_name}.V"].T
    ) * normalization_factor
    • Here, mm corresponds to m\mathbf{m}.
    • decomposed_params[f"{param_name}.S"] * mm performs element-wise multiplication σi×mi\sigma_i \times m_i.
    • The expression reconstructs W′=U⋅diag(σi×mi)⋅VTW' = U \cdot \text{diag}(\sigma_i \times m_i) \cdot V^T
  • forward Function (utils.py):
    for k in base_params:
    if "mlp" in k:
    new_params[k] = compose_new_params(policy, k, decomposed_params, learnable_params)
    model.get_parameter(k).copy_(new_params[k])
    else:
    new_params[k] = base_params[k]
    • For each MLP-related weight kk, new_params[k] represents the updated weight W′W' for that layer.

Understanding the Mask

A mask m\mathbf{m} is a vector of scaling factors applied to singular values:

σi′=σi×mi\sigma_i' = \sigma_i \times m_i

In code, the mask mm modulates singular values as shown:

mm = policy.get_mask(learnable_params[param_name])

Here, mm[i] corresponds to mim_i scaling the ii-th singular value σi\sigma_i.

Multi-Layer Perceptron (MLP)

An MLP (Multi-Layer Perceptron) is a feed-forward neural network:

  • It consists of layers of interconnected neurons.
  • In transformer models, the MLP refers to the feed-forward portion in each layer following the self-attention mechanism.

When the code checks:

if "mlp" in k:
...

it targets parameters belonging to MLP layers, indicating that the SVF adjustments are applied specifically to these dense layers.


2. Reinforcement Learning Policy Gradient

Mathematical Concept:

Policy gradient methods update parameters θ\theta to maximize expected rewards:

∇θJ(θ)≈E[−log⁡p(y∣x;θ)⋅r]\nabla_{\theta} J(\theta) \approx \mathbb{E}[-\log p(y|x; \theta) \cdot r]

with optional KL-divergence regularization:

loss=−log⁡p(y∣x;θ)⋅r+λDKL(pθ∥pref)\text{loss} = -\log p(y|x; \theta) \cdot r + \lambda D_{\text{KL}}(p_{\theta} \| p_{\text{ref}})

Python-to-Math Variable Correspondence:

  • xx and yy: Input prompt and generated output text.
  • θ\theta: Policy parameters updated during RL, corresponding to weights in objects like policy.
  • rr: Reward signal, computed from correctness of the model’s output.
  • α\alpha: In weighted combinations, these are the adaptive coefficients stored as adaptive_weights in the policy.

Code Implementation in Reinforce class:

  • Policy Gradient Calculation:
    log_likelihood = selected_log_probs.sum(axis=-1)
    pg = -log_likelihood * rewards[j]
    loss = pg
    if use_kl_loss:
    kl_div = F.kl_div(...)
    loss = loss + kl_ref_coeff * kl_div
    scaled_loss = loss / clipped_batch_size
    scaled_loss.backward()
    • log_likelihood computes log⁡p(y∣x;θ)\log p(y|x; \theta).
    • pg = -log_likelihood * rewards[j] corresponds to −log⁡p(y∣x;θ)⋅r-\log p(y|x; \theta) \cdot r.
    • KL divergence computation and addition mirror the regularized term λD_KL(⋅)\lambda D\_{\text{KL}}(\cdot).
  • Parameter Update:
    def update(self, policy):
    torch.nn.utils.clip_grad_norm_(policy.trainable_params, max_grad_norm)
    self.optimizer.step()
    self.optimizer.zero_grad()
    • This updates the RL policy parameters θ\theta using gradients derived from the computed loss.

3. Weighted Combination of Expert Models

Mathematical Concept:

Weighted combination uses coefficients α\alpha to blend expert weights:

Wcombined=∑i=1Nαi W(i)W_{\text{combined}} = \sum_{i=1}^{N} \alpha_i \, W^{(i)}

where:

  • W(i)W^{(i)} are weight matrices from expert model ii.
  • αi\alpha_i are combination coefficients, analogous to variables in code like adaptive_weights.

Python-to-Math Variable Correspondence:

  • adaptive_weights: Corresponds to coefficients α\alpha.
  • vs: A list containing expert weights W(i)W^{(i)} for a specific parameter kk.
  • output_params[k]: Represents the combined weight W_combinedW\_{\text{combined}} for parameter kk.

Code Implementation in WeightedCombination class:

  • Combining Weights:
    def get_learnable_params(self):
    adaptive_coeff_per_layer = self.get_coeff_per_layer()
    output_params = {}
    for i, (k, vs) in enumerate(self.original_params.items()):
    cs_coeff = adaptive_coeff_per_layer[:, i]
    out = vs[0] * cs_coeff[0]
    for j, other_v in enumerate(vs[1:]):
    out = out + other_v * cs_coeff[j+1]
    output_params[k] = out
    return output_params
    • For each parameter key kk, this computes: Wk=∑i=1Nαi,k vk(i)W_k = \sum_{i=1}^{N} \alpha_{i,k} \, v^{(i)}_k where cs_coeff[j] corresponds to α_j,k\alpha\_{j,k} and each vs[j] corresponds to W(j)_kW^{(j)}\_k.
    • The result output_params[k] is the weighted combination W_combinedW\_{\text{combined}} for that parameter.

4. Summary of Variable Correspondences

  • W,W′W, W' (e.g., base_params[k], new_params[k]): Weight matrices before and after fine-tuning.
  • U,S,VU, S, V (e.g., decomposed_params[...]): Matrices from SVD decomposition used in SVF.
  • σi\sigma_i (e.g., elements of decomposed_params[f"{param_name}.S"]): Singular values of weight matrices.
  • m\mathbf{m} or mask (e.g., mm): Scaling factors applied to singular values during fine-tuning.
  • θ\theta (implied in policy parameters): Parameters of the policy being optimized via reinforcement learning.
  • rr (e.g., rewards[j]): Reward signal for a given output, derived from task correctness.
  • α\alpha (e.g., adaptive_weights, adaptive_coeff_per_layer): Coefficients for weighted combination of expert models.
  • x,yx, y (e.g., prompt, result.generation): Input prompts and generated outputs during evaluation.

Conclusion

The three pieces fit together: SVD decomposes each MLP weight matrix, a learned mask m\mathbf{m} scales its singular values, and a policy gradient trains that mask against a task reward rr. A separate weighted-combination mode blends several fine-tuned experts with coefficients α\alpha instead of using a single mask.

Mapping the code variables (mm, adaptive_weights, new_params) to the math (m\mathbf{m}, α\alpha, W′W') is the fastest way to read the rest of the repository, since the paper’s notation and the implementation track each other closely.

Adapting a base model’s weights at inference time is only half the production problem: you also need to know when the adapted weights are actually helping. The tooling below covers that observability layer, and the newsletter is where write-ups like this one ship first.

Next steps: scaling to production

If you take this into production, these are the pieces I would add first.

  • Supabase Supabase is a hosted Postgres platform with authentication and storage built in. Postgres with pgvector for embeddings, so you do not run a separate vector store.
  • Datadog Datadog aggregates metrics, logs, and traces for infrastructure monitoring. Traces and cost metrics across model calls, so latency and spend are visible per request.
  • Vercel Vercel hosts frontend applications with a global edge network and CI/CD. Deploys the frontend and edge functions that sit in front of the model API.

Deploying generative AI models to production

Get the free playbook on shipping generative AI models to production.