Memory-Efficient Model Weight Loading Strategies in LLMs-from-Scratch

Use torch.load() with map_location, prefer the safetensors format, and load sharded checkpoints parameter-by-parameter to keep RAM usage below single-GPU limits when initializing large language models.

The rasbt/LLMs-from-scratch repository demonstrates how to load billion-parameter model weights without exhausting system memory. By combining device-aware streaming, memory-mapped file formats, and selective parameter loading, the codebase enables practitioners to initialize models that far exceed available RAM. These techniques are implemented across the pkg/llms_from_scratch module and illustrated in the chapter-specific notebooks.

Device-Aware Lazy Loading

The foundation of memory-efficient loading in this repository is device-aware checkpoint streaming. Instead of loading weights into CPU RAM before transferring to GPU, the code uses torch.load(..., map_location=device) to stream tensors directly onto the target device.

This approach eliminates the intermediate copy that typically doubles memory consumption during load operations. In pkg/llms_from_scratch/utils.py, helper functions wrap this pattern to ensure all loading utilities follow the same memory-efficient signature.

import torch
from llms_from_scratch.ch05 import load_weights_into_gpt

# Stream directly to CPU (or "cuda" if available) without intermediate copy

params = torch.load("gpt2_weights.pt", map_location="cpu", weights_only=True)
model = GPTModel(cfg)
load_weights_into_gpt(model, params)  # Copies only required tensors

Sharded Checkpoints and Safetensors Format

For models too large to store in a single file, the repository implements sharded checkpoint loading using the safetensors format. Unlike PyTorch's native pickle-based .pt files, safetensors provides a binary-only, memory-safe specification that supports memory mapping.

In pkg/llms_from_scratch/qwen3.py (lines 598-699), the download_from_huggingface function retrieves model manifests and iterates over shard files. Each shard is loaded using safetensors.torch.load_file, which maps the file directly into tensors without Python deserialization overhead.

from llms_from_scratch.qwen3 import download_from_huggingface, load_weights_into_qwen
import torch

# Download sharded weights (lines 598-699 in qwen3.py)

download_from_huggingface(repo_id="Qwen/Qwen-3-14B", local_dir="qwen3_weights")

# Load specific shards with memory mapping (lines 673-685)

weights = torch.load("qwen3_weights/model-00001-of-00004.safetensors", map_location="cpu")
model = Qwen3Model(cfg)
load_weights_into_qwen(model, cfg, weights)  # Defined at line 452-456

Parameter-Wise Mapping Strategy

The load_weights_into_* family of functions implements selective parameter loading to skip unnecessary tensors. Each loader accepts a param_config dictionary that maps checkpoint keys to specific model sub-modules.

Rather than loading the entire state dict into memory, these functions iterate only over required keys. This implementation appears in pkg/llms_from_scratch/ch05.py for GPT models and is mirrored in pkg/llms_from_scratch/llama3.py (line 567) and pkg/llms_from_scratch/qwen3.py for their respective architectures.


# From ch05.py - load_weights_into_gpt demonstrates selective key iteration

def load_weights_into_gpt(model, params):
    # Only process keys that exist in the target model configuration

    for name, param in model.named_parameters():
        if name in params:
            param.data.copy_(params[name])

Unified Regex-Based Refactoring

To ensure consistency across the codebase, pkg/llms_from_scratch/utils.py (lines 100-104) contains a regular-expression rewrite that standardizes loader signatures. This utility automatically converts legacy function definitions that accepted class names into the modern signature expecting a model instance.


# utils.py pattern standardizes loader signatures

pattern = r"(def\s+load_weights_into_\w+\s*\()\s*\w+\s*,"
new_pat = r"\1model,"

This refactoring guarantees that all weight-loading helpers follow the memory-efficient pattern of accepting a model instance and a parameter dictionary, rather than instantiating classes internally.

Summary

  • Device-aware streaming via torch.load(map_location=...) eliminates intermediate CPU copies when loading to GPU.
  • Safetensors format in pkg/llms_from_scratch/qwen3.py enables memory-mapped file access without pickle serialization overhead.
  • Sharded checkpoints allow models to be split across multiple files, loading only the necessary shards for a given operation.
  • Parameter-wise loading through load_weights_into_gpt and similar functions iterates only over required keys, skipping unrelated tensors.
  • Regex standardization in utils.py ensures all loader functions follow the same memory-efficient signature pattern.

Frequently Asked Questions

How does device-aware loading reduce memory usage?

Device-aware loading passes the map_location argument to torch.load(), which streams tensors directly onto the target device (CPU or CUDA) rather than loading into system RAM first and then copying. This avoids holding two copies of the weights simultaneously—the original load and the device transfer—effectively halving peak memory consumption during initialization.

What is the advantage of using safetensors over PyTorch's native format?

The safetensors format provides a binary-only, memory-safe specification that supports direct memory mapping, whereas PyTorch's .pt format relies on Python pickle deserialization which requires loading the entire file into memory first. According to the implementation in pkg/llms_from_scratch/qwen3.py (lines 673-685), safetensors.torch.load_file can map file contents directly into tensors without temporary copies, significantly reducing RAM requirements for large checkpoints.

How does parameter-wise mapping help with large models?

Parameter-wise mapping allows the load_weights_into_* functions to iterate only over the specific keys required by the model configuration rather than loading the entire state dictionary. As implemented in pkg/llms_from_scratch/ch05.py, this selective approach skips unrelated tensors in the checkpoint file, saving both I/O bandwidth and memory when loading partial models or adapter weights.

Where can I find the memory-efficient loading implementation for GPT models?

The GPT-specific implementation resides in pkg/llms_from_scratch/ch05.py within the load_weights_into_gpt function. Practical usage examples are demonstrated in ch05/01_main-chapter-code/gpt_generate.py, which shows how to combine the loader with the GPTModel class defined in ch05/08_memory_efficient_weight_loading/previous_chapters.py for text generation tasks.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →