I read the paper and code so you don't have to: Self-Adaptive Large Language Models
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:
where contains singular values .
-
Fine-tuning with a Mask: Adjust the singular values with a learnable vector (or mask ) to obtain new weights:
The mask scales each singular value, modulating the contribution of each singular component in .
Python-to-Math Variable Correspondence:
Win the code corresponds to a weight matrix of the model.U,S, andV(extracted asdecomposed_params[f"{param_name}.U"],...S,...Vin the code) represent the components from the SVD of .mm(frompolicy.get_mask(...)) is the mask .- The resulting recomposed matrix
W'is stored in variables likenew_params[k].
Code Implementation:
compose_new_paramsFunction (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,
mmcorresponds to . decomposed_params[f"{param_name}.S"] * mmperforms element-wise multiplication .- The expression reconstructs
- Here,
forwardFunction (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 ,
new_params[k]represents the updated weight for that layer.
- For each MLP-related weight ,
Understanding the Mask
A mask is a vector of scaling factors applied to singular values:
In code, the mask mm modulates singular values as shown:
mm = policy.get_mask(learnable_params[param_name])Here, mm[i] corresponds to scaling the -th singular value .
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 to maximize expected rewards:
with optional KL-divergence regularization:
Python-to-Math Variable Correspondence:
- and : Input prompt and generated output text.
- : Policy parameters updated during RL, corresponding to weights in objects like
policy. - : Reward signal, computed from correctness of the model’s output.
- : In weighted combinations, these are the adaptive coefficients stored as
adaptive_weightsin the policy.
Code Implementation in Reinforce class:
- Policy Gradient Calculation:
log_likelihood = selected_log_probs.sum(axis=-1)pg = -log_likelihood * rewards[j]loss = pgif use_kl_loss:kl_div = F.kl_div(...)loss = loss + kl_ref_coeff * kl_divscaled_loss = loss / clipped_batch_sizescaled_loss.backward()
log_likelihoodcomputes .pg = -log_likelihood * rewards[j]corresponds to .- KL divergence computation and addition mirror the regularized term .
- 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 using gradients derived from the computed loss.
3. Weighted Combination of Expert Models
Mathematical Concept:
Weighted combination uses coefficients to blend expert weights:
where:
- are weight matrices from expert model .
- are combination coefficients, analogous to variables in code like
adaptive_weights.
Python-to-Math Variable Correspondence:
adaptive_weights: Corresponds to coefficients .vs: A list containing expert weights for a specific parameter .output_params[k]: Represents the combined weight for parameter .
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] = outreturn output_params
- For each parameter key , this computes:
where
cs_coeff[j]corresponds to and eachvs[j]corresponds to . - The result
output_params[k]is the weighted combination for that parameter.
- For each parameter key , this computes:
where
4. Summary of Variable Correspondences
- (e.g.,
base_params[k],new_params[k]): Weight matrices before and after fine-tuning. - (e.g.,
decomposed_params[...]): Matrices from SVD decomposition used in SVF. - (e.g., elements of
decomposed_params[f"{param_name}.S"]): Singular values of weight matrices. - or mask (e.g.,
mm): Scaling factors applied to singular values during fine-tuning. - (implied in policy parameters): Parameters of the policy being optimized via reinforcement learning.
- (e.g.,
rewards[j]): Reward signal for a given output, derived from task correctness. - (e.g.,
adaptive_weights,adaptive_coeff_per_layer): Coefficients for weighted combination of expert models. - (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 scales its singular values, and a policy gradient trains that mask against a task reward . A separate weighted-combination mode blends several fine-tuned experts with coefficients instead of using a single mask.
Mapping the code variables (mm, adaptive_weights, new_params) to the math (, , ) 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.