Layout¶
layout
¶
channels_last
¶
Context manager that converts a tensor to channels-last on entry and back to channels-first on exit.
Usage::
ref = [tensor_nchw]
with channels_last(ref):
# ref[0] is now in NHWC format
result = some_mlx_op(ref[0])
ref[0] = result
# ref[0] is back in NCHW format
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x_ref
|
list[array]
|
Single-element list containing the tensor. Modified in-place. |
required |
to_channels_first
¶
Convert from channels-last to channels-first format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Tensor in channels-last format. 3D: (B, L, C) -> (B, C, L) 4D: (B, H, W, C) -> (B, C, H, W) 5D: (B, D, H, W, C) -> (B, C, D, H, W) |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tensor in channels-first format. |
to_channels_last
¶
Convert from channels-first to channels-last format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Tensor in channels-first format. 3D: (B, C, L) -> (B, L, C) 4D: (B, C, H, W) -> (B, H, W, C) 5D: (B, C, D, H, W) -> (B, D, H, W, C) |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tensor in channels-last format. |
convert_conv_weights
¶
Convert a convolution weight tensor from PyTorch to MLX format.
PyTorch conv weights are channels-first: (out, in, *kernel). MLX
keeps the output channel first and moves the input channel last:
- Conv1d:
(out, in, K)→(out, K, in) - Conv2d:
(out, in, kH, kW)→(out, kH, kW, in) - Conv3d:
(out, in, kD, kH, kW)→(out, kD, kH, kW, in)
This function handles the permutation for all conv dimensions.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
weight
|
array
|
PyTorch-format conv weight tensor. |
required |
Returns:
| Type | Description |
|---|---|
array
|
MLX-format conv weight tensor. |
load_safetensors
¶
load_safetensors(path: str, key_map: dict[str, str] | None = None, key_fn: Callable[[str], str] | None = None, conv_keys: set[str] | None = None) -> dict[str, array]
Load safetensors weights with optional key remapping and conv conversion.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
path
|
str
|
Path to .safetensors file. |
required |
key_map
|
dict[str, str] | None
|
Optional dict mapping source keys to target keys. Keys not in the map are kept as-is. |
None
|
key_fn
|
Callable[[str], str] | None
|
Optional function to transform key names. Applied after key_map. |
None
|
conv_keys
|
set[str] | None
|
Set of key names (after remapping) that contain convolution weights and should be permuted from PyTorch to MLX format. |
None
|
Returns:
| Type | Description |
|---|---|
dict[str, array]
|
Dict of parameter name -> mx.array. |
channels
¶
Channel layout conversion utilities.
MLX uses channels-last (NHWC/NDHWC) while PyTorch uses channels-first (NCHW/NCDHW). These utilities handle the conversion.
to_channels_last
¶
Convert from channels-first to channels-last format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Tensor in channels-first format. 3D: (B, C, L) -> (B, L, C) 4D: (B, C, H, W) -> (B, H, W, C) 5D: (B, C, D, H, W) -> (B, D, H, W, C) |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tensor in channels-last format. |
to_channels_first
¶
Convert from channels-last to channels-first format.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
array
|
Tensor in channels-last format. 3D: (B, L, C) -> (B, C, L) 4D: (B, H, W, C) -> (B, C, H, W) 5D: (B, D, H, W, C) -> (B, C, D, H, W) |
required |
Returns:
| Type | Description |
|---|---|
array
|
Tensor in channels-first format. |
channels_last
¶
Context manager that converts a tensor to channels-last on entry and back to channels-first on exit.
Usage::
ref = [tensor_nchw]
with channels_last(ref):
# ref[0] is now in NHWC format
result = some_mlx_op(ref[0])
ref[0] = result
# ref[0] is back in NCHW format
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x_ref
|
list[array]
|
Single-element list containing the tensor. Modified in-place. |
required |
weights
¶
Weight conversion utilities for loading PyTorch models into MLX.
convert_conv_weights
¶
Convert a convolution weight tensor from PyTorch to MLX format.
PyTorch conv weights are channels-first: (out, in, *kernel). MLX
keeps the output channel first and moves the input channel last:
- Conv1d:
(out, in, K)→(out, K, in) - Conv2d:
(out, in, kH, kW)→(out, kH, kW, in) - Conv3d:
(out, in, kD, kH, kW)→(out, kD, kH, kW, in)
This function handles the permutation for all conv dimensions.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
weight
|
array
|
PyTorch-format conv weight tensor. |
required |
Returns:
| Type | Description |
|---|---|
array
|
MLX-format conv weight tensor. |
load_safetensors
¶
load_safetensors(path: str, key_map: dict[str, str] | None = None, key_fn: Callable[[str], str] | None = None, conv_keys: set[str] | None = None) -> dict[str, array]
Load safetensors weights with optional key remapping and conv conversion.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
path
|
str
|
Path to .safetensors file. |
required |
key_map
|
dict[str, str] | None
|
Optional dict mapping source keys to target keys. Keys not in the map are kept as-is. |
None
|
key_fn
|
Callable[[str], str] | None
|
Optional function to transform key names. Applied after key_map. |
None
|
conv_keys
|
set[str] | None
|
Set of key names (after remapping) that contain convolution weights and should be permuted from PyTorch to MLX format. |
None
|
Returns:
| Type | Description |
|---|---|
dict[str, array]
|
Dict of parameter name -> mx.array. |