Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions .github/workflows/upstream-sync.yml
Original file line number Diff line number Diff line change
Expand Up @@ -32,13 +32,14 @@ jobs:

- name: Create or Update Sync Branch
run: |
git checkout -B sync/upstream-latest
git reset --hard upstream/main
git checkout -B sync/upstream-latest origin/main
git checkout upstream/main -- .
git checkout origin/main -- .github/workflows
git push -f origin sync/upstream-latest

- name: Create Pull Request to Main
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
Comment on lines 42 to 44
gh pr create --base main --head sync/upstream-latest \
--title "🔄 Auto-Sync: Apple Upstream Repository" \
Expand Down
19 changes: 19 additions & 0 deletions Source/MLX/MLXFast.swift
Original file line number Diff line number Diff line change
Expand Up @@ -420,6 +420,25 @@ public enum MLXFast {
mlx_fast_set_prefetch_enabled(enabled)
}

/// Convert an array of raw FP8 E4M3 bytes (stored as uint8 by safetensors loader)
/// to the specified floating point dtype using proper FP8 E4M3 semantics.
///
/// MLX's safetensors loader maps `F8_E4M3` → `uint8` (raw bit patterns).
/// Use this before applying block-wise scale_inv to dequantize FP8 weights correctly.
///
/// - Parameters:
/// - x: uint8 array containing raw FP8 E4M3 bit patterns
/// - dtype: target floating-point dtype (e.g. `.bfloat16`, `.float16`, `.float32`)
/// - stream: stream or device to evaluate on
/// - Returns: Array in the requested dtype with correctly converted FP8 values
public static func fromFp8(
_ x: MLXArray, dtype: DType = .bfloat16, stream: StreamOrDevice = .default
) -> MLXArray {
var result = mlx_array_new()
mlx_from_fp8(&result, x.ctx, dtype.cmlxDtype, stream.ctx)
return MLXArray(result)
Comment on lines +434 to +441
}

}

/// Optimized implementation of `NN.RoPE`.
Expand Down
Loading