flyte.extras.model_streamer
Stream safetensors model weights from object storage straight onto a device.
Weights are fetched as parallel byte ranges and handed over tensor by tensor as each one completes. They never touch local disk, and host memory holds only the tensors in flight, so a model’s load time is close to its download time.
ModelStreamer: the streamer itself (async and sync iterators, a state dict, or straight into annn.Module).load_hf_model: atransformersmodel built onmetaand filled on the GPU.flyteplugins.vllm.model_streamer(inflyteplugins-vllm): a vLLMload_formatforvllm.LLM/AsyncLLMEngine.
Requires torch; load_hf_model also needs transformers, and the
vLLM integration needs vllm.
Directory
Classes
| Class | Description |
|---|---|
LoadResult |
What ModelStreamer.load_into did not match. |
ModelStreamer |
Stream the safetensors weights under a remote prefix, tensor by tensor. |
Methods
| Method | Description |
|---|---|
empty_weights() |
Build modules with their parameters on the meta device. |
load_hf_model() |
Build a transformers model and stream its weights onto device. |
missing_parameters() |
Parameters and buffers still on the meta device. |
Methods
empty_weights()
def empty_weights()Build modules with their parameters on the meta device.
Unlike with torch.device("meta"), buffers are created normally, so
non-persistent buffers computed in __init__ (rotary inv_freq, for
instance), which are never in a checkpoint, keep real values.
load_hf_model()
def load_hf_model(
path: str,
device: torch.device | str = 'cuda',
dtype: torch.dtype | None = None,
model_class: typing.Any = None,
local_dir: str | pathlib.Path | None = None,
trust_remote_code: bool = False,
chunk_size: int | None = None,
max_concurrency: int | None = None,
) -> tuple[nn.Module, pathlib.Path]Build a transformers model and stream its weights onto device.
The config, tokenizer and generation config are downloaded to
local_dir (a fresh temporary directory by default). The model is built
with its parameters on meta, and each weight is then copied onto
device as soon as it finishes downloading. Weights never touch local
disk, and host memory holds only the tensors in flight.
model_class defaults to AutoModelForCausalLM; any Auto* class
or concrete PreTrainedModel subclass works. dtype casts
floating-point weights (None keeps the checkpoint’s dtype).
Returns the model in eval mode and the local directory, from which the
tokenizer can be loaded with AutoTokenizer.from_pretrained(local_dir).
| Parameter | Type | Description |
|---|---|---|
path |
str |
|
device |
torch.device | str |
|
dtype |
torch.dtype | None |
|
model_class |
typing.Any |
|
local_dir |
str | pathlib.Path | None |
|
trust_remote_code |
bool |
|
chunk_size |
int | None |
|
max_concurrency |
int | None |
missing_parameters()
def missing_parameters(
module: nn.Module,
) -> list[str]Parameters and buffers still on the meta device.
| Parameter | Type | Description |
|---|---|---|
module |
nn.Module |