diff --git a/.cursor/prompts/create-pr-description.md b/.cursor/prompts/create-pr-description.md new file mode 100644 index 0000000..80b84c9 --- /dev/null +++ b/.cursor/prompts/create-pr-description.md @@ -0,0 +1,79 @@ +# Create Pull Request Description + +Generate a concise, benefit-focused PR description in US English. + +## Guidelines + +### Structure + +```markdown +## Summary + +[One paragraph explaining WHAT changed and WHY it matters] + +## Key Changes + +- [Bullet points of significant changes - focus on impact, not implementation details] + +## Architecture Decisions + +[Only include if there are decisions other contributors should be aware of] + +## Breaking Changes + +[Only include if there are breaking changes] +``` + +### Writing Style + +- **Be concise**: Every sentence should add value +- **Focus on benefits**: What does this enable? What problem does it solve? +- **Avoid redundancy**: Don't repeat information, don't state the obvious +- **Skip boilerplate**: No "This PR adds...", no test mentions (CI handles that) +- **Use active voice**: "Ports X from Y" not "X was ported from Y" + +### What to Include + +- Significant architectural changes +- New capabilities or features +- Performance improvements with context +- Migration guidance if needed +- Links to related issues/RFCs + +### What to Exclude + +- Test coverage details (CI shows this) +- Obvious file changes (reviewers can see the diff) +- Implementation minutiae +- Changelog-style lists of every file touched + +## Example + +```markdown +## Summary + +Switches MLX infrastructure from vendored mlx-swift-lm to direct ports from mlx-lm (Python). This gives us access to the latest model architectures faster, as mlx-lm releases more frequently and has broader model coverage. + +## Key Changes + +- Direct Python→Swift ports for KVCache, RoPE, and MoE layers +- New `ported/` directory structure with version tracking +- Generator now produces code matching mlx-swift-lm patterns exactly + +## Architecture Decisions + +**Why port from Python instead of using mlx-swift-lm?** +mlx-lm (Python) is the primary source, updated more frequently, and supports models like Llama 4 MoE that mlx-swift-lm doesn't yet have. + +**Directory structure**: + +- `generated/models/` - hf2swift generator output +- `ported/` - LLM-assisted ports from Python with git hash tracking +``` + +## Instructions + +1. Analyze the current branch changes using `git log` and `git diff` +2. Read any relevant RFCs or decision documents +3. Generate a PR description following the structure above +4. Keep total length under 500 words diff --git a/.cursor/prompts/port-python-to-swift.md b/.cursor/prompts/port-python-to-swift.md new file mode 100644 index 0000000..8d10981 --- /dev/null +++ b/.cursor/prompts/port-python-to-swift.md @@ -0,0 +1,243 @@ +# Port Python mlx-lm to Swift + +You are porting Python code from Apple's `mlx-lm` library to Swift for the `node-mlx` project. + +## Source Repository + +- **Primary**: https://github.com/ml-explore/mlx-lm/tree/main/mlx_lm/models +- **Reference only**: https://github.com/ml-explore/mlx-swift-lm + +**IMPORTANT**: Always record the exact git hash. Get it with: + +```bash +curl -s "https://api.github.com/repos/ml-explore/mlx-lm/commits/main" | grep '"sha"' | head -1 +``` + +## File Locations + +| Type | Directory | +| ----------------- | ------------------------------------------------------ | +| Ported code | `packages/swift/Sources/NodeMLXCore/ported/` | +| Shared components | `packages/swift/Sources/NodeMLXCore/shared/` | +| Tests | `packages/swift/Tests/NodeMLXCoreTests/` | +| Generated models | `packages/swift/Sources/NodeMLXCore/generated/models/` | + +## Core Principles + +### 1. Clean Cut Philosophy + +- Start fresh, don't patch existing code +- Port with understanding, not blind translation +- Premium architect-level Swift: idiomatic, elegant, maintainable + +### 2. Focus on Popular Models + +| Priority | Models | Notes | +| ------------ | ----------------------------------------- | ------------------------- | +| ✅ Essential | Llama, Qwen, Phi, Gemma, Mistral, GPT-OSS | Mainstream | +| ⏸️ Defer | Mamba, Jamba, DBRX | SSM/unusual architectures | +| ❌ Skip | Batch processing, server features | Not needed for inference | + +### 3. Minimal Viable Port + +- Port core functionality, not edge cases +- Skip features that < 5% of users need +- Add extensibility points for future additions + +## File Header Template + +Every ported file **must** include: + +```swift +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/.py +// Git Hash: () +``` + +## Swift Style Guide + +### Naming Conventions + +| Python | Swift | +| ---------------------- | --------------------- | +| `snake_case` | `camelCase` | +| `class KVCache` | `class KVCache` | +| `def update_and_fetch` | `func updateAndFetch` | +| `__init__` | `init` | +| `__len__` | `var count: Int` | +| `_private_method` | `private func method` | + +### Type Mappings + +| Python | Swift | +| ------------- | ----------------- | +| `mx.array` | `MLXArray` | +| `nn.Module` | `Module` (MLXNN) | +| `Optional[T]` | `T?` | +| `List[T]` | `[T]` | +| `Dict[K, V]` | `[K: V]` | +| `Tuple[A, B]` | `(A, B)` | +| `None` | `nil` | +| `@property` | computed property | + +### MLX Operations + +| Python | Swift | +| -------------------------------- | ------------------------------- | +| `mx.zeros(shape)` | `MLXArray.zeros(shape)` | +| `mx.concatenate([a, b], axis=2)` | `concatenated([a, b], axis: 2)` | +| `mx.quantize(x, ...)` | `MLX.quantized(x, ...)` | +| `x[..., :n, :]` | `x[.ellipsis, .. (MLXArray, MLXArray) + var offset: Int { get } + func makeMask(queryLength: Int, windowSize: Int?) -> MLXFast.ScaledDotProductAttentionMaskMode +} +``` + +### Class Structure + +```swift +/// KV cache with grow-in-place strategy +public class StandardKVCache: KVCacheProtocol { + // MARK: - Properties + + private var keys: MLXArray? + private var values: MLXArray? + public private(set) var offset: Int = 0 + + public static let step = 256 + + // MARK: - Initialization + + public init() {} + + // MARK: - Cache Operations + + public func update(keys: MLXArray, values: MLXArray) -> (MLXArray, MLXArray) { + // Implementation + } +} +``` + +## What NOT to Port + +### From cache.py + +- ❌ `BatchKVCache`, `BatchRotatingKVCache` - Server/batch processing +- ❌ `MambaCache`, `ArraysCache` - SSM models +- ❌ `ChunkedKVCache`, `CacheList` - Specialized use cases +- ❌ `save_prompt_cache`, `load_prompt_cache` - Serialization + +### General + +- ❌ Batch processing features +- ❌ Prompt caching to disk +- ❌ Speculative decoding caches +- ❌ Multi-modal (initially) + +## Shared Components + +Before porting, check if a shared component already exists in `shared/`: + +| Component | File | Use When | +| ----------------- | --------------------------- | -------------------------- | +| RMSNorm | `RMSNorm.swift` | Standard RMS normalization | +| GemmaRMSNorm | `ported/GemmaRMSNorm.swift` | (1+weight) scaling | +| StandardAttention | `StandardAttention.swift` | Basic GQA attention | +| StandardMLP | `StandardMLP.swift` | SwiGLU MLP | +| MathUtils | `MathUtils.swift` | erfinv, clipResidual, topK | + +## Testing + +### Test File Location + +Tests go in `packages/swift/Tests/NodeMLXCoreTests/`: + +```swift +import XCTest +@testable import NodeMLXCore +import MLX + +final class KVCacheTests: XCTestCase { + func testUpdateAndFetch() { + let cache = StandardKVCache() + let keys = MLXArray.zeros([1, 4, 8, 64]) + let values = MLXArray.zeros([1, 4, 8, 64]) + + let (k, v) = cache.update(keys: keys, values: values) + + XCTAssertEqual(cache.offset, 8) + XCTAssertEqual(k.dim(2), 8) + } +} +``` + +### Running Tests + +```bash +cd packages/swift +swift test +``` + +## Workflow + +1. **Download Python source**: + + ```bash + curl -s "https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/.py" -o /tmp/.py + ``` + +2. **Analyze**: Essential vs. optional features + +3. **Check shared components**: Reuse if exists + +4. **Design Swift API**: Protocols, classes + +5. **Implement**: Premium Swift patterns + +6. **Test**: Comprehensive coverage + +7. **Document**: Update PORTING_DECISIONS.md + +8. **Build**: + ```bash + cd packages/swift && swift build -c release && swift test + ``` + +## Documentation Updates + +After porting, update: + +1. **File header**: Git hash, date +2. **PORTING_DECISIONS.md**: What was ported, decisions made +3. **ported/README.md**: Add to ported files table +4. **Tests**: Add test file + +## Quick Reference + +```bash +# Get latest mlx-lm hash +curl -s "https://api.github.com/repos/ml-explore/mlx-lm/commits/main" | grep '"sha"' | head -1 + +# Download Python source +curl -s "https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/cache.py" -o /tmp/cache.py + +# Build and test +cd packages/swift && swift build -c release && swift test + +# Regenerate models (to ensure compatibility) +pnpm hf2swift --model llama --output packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift +``` diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 411d209..cc433f7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -39,9 +39,9 @@ jobs: uses: actions/cache@v4 with: path: packages/swift/.build - key: swift-build-${{ runner.os }}-${{ hashFiles('packages/swift/Package.resolved', 'packages/swift/Package.swift') }} + key: swift-build-v3-${{ runner.os }}-${{ hashFiles('packages/swift/Package.resolved', 'packages/swift/Package.swift', 'packages/swift/Sources/**/*.swift') }} restore-keys: | - swift-build-${{ runner.os }}- + swift-build-v3-${{ runner.os }}- - name: Cache Xcode DerivedData uses: actions/cache@v4 @@ -73,32 +73,29 @@ jobs: env: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} - # Swift Tests with Coverage - - name: Run Swift tests with coverage + # Swift Tests (unit tests - no model downloads) + - name: Run Swift unit tests working-directory: ./packages/swift run: | - xcodebuild test \ - -scheme NodeMLX \ - -destination 'platform=macOS' \ - -enableCodeCoverage YES \ - -resultBundlePath ./test-results.xcresult \ - 2>&1 | xcbeautify || true - - - name: Export Swift coverage - working-directory: ./packages/swift - run: | - # Convert xcresult to JSON format for Codecov - xcrun xccov view --report --json test-results.xcresult > coverage.json || true - - - name: Upload Swift coverage to Codecov - uses: codecov/codecov-action@v5 - with: - files: ./packages/swift/coverage.json - flags: swift - name: swift-coverage - fail_ci_if_error: false - env: - CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }} + # Build tests with testing enabled (builds everything including library) + swift build -c release -Xswiftc -enable-testing --build-tests + + # Copy Metal library to test bundle location + TEST_BUNDLE=".build/arm64-apple-macosx/release/NodeMLXPackageTests.xctest/Contents/MacOS" + METALLIB=".build/arm64-apple-macosx/release/mlx-swift_Cmlx.bundle/Contents/Resources/default.metallib" + if [ -f "$METALLIB" ]; then + mkdir -p "$TEST_BUNDLE" + cp "$METALLIB" "$TEST_BUNDLE/mlx.metallib" + echo "✓ Copied mlx.metallib to test bundle" + else + echo "⚠ mlx.metallib not found at $METALLIB - tests may fail" + find .build -name "*.metallib" 2>/dev/null || true + fi + + # Run unit tests only (skip integration tests that require model downloads) + # Integration tests run via the smoke test below with cached models + # Filter pattern: list all unit test classes explicitly + swift test -c release --skip-build --filter 'GenerateTests|KVCacheTests|MaskTests|PerformanceTests|RoPEUtilsTests|SamplingUtilsTests|StringOrNumberTests|SwitchLayersTests' - name: Verify Swift library run: test -f packages/node-mlx/swift/libNodeMLX.dylib diff --git a/.gitignore b/.gitignore index 8fe72c9..7bdad31 100644 --- a/.gitignore +++ b/.gitignore @@ -47,3 +47,4 @@ DerivedData/ # Temporary files *.tmp *.bak +.venv/ diff --git a/.husky/pre-push b/.husky/pre-push index ff90472..d80bbcc 100755 --- a/.husky/pre-push +++ b/.husky/pre-push @@ -16,8 +16,59 @@ pnpm typecheck || { exit 1 } +# Regenerate Swift models and check for uncommitted changes +echo "→ Regenerating Swift models..." +MODELS_DIR="packages/swift/Sources/NodeMLXCore/generated/models" + +# List of models that ARE auto-generated +GENERATED_MODELS=( + "qwen2:Qwen2Generated.swift" + "qwen3:Qwen3Generated.swift" + "llama:LlamaGenerated.swift" + "phi3:Phi3Generated.swift" + "gemma3:Gemma3Generated.swift" + "gemma3n:Gemma3nGenerated.swift" + "mistral:MistralGenerated.swift" + "mistral3:Mistral3Generated.swift" + "smollm3:SmolLM3Generated.swift" + "gpt_oss:GptOSSGenerated.swift" +) + +# Build hf2swift first +pnpm --filter @node-mlx/hf2swift build --silent 2>/dev/null || { + echo "⚠️ Could not build hf2swift, skipping model regeneration check" +} + +# Check if hf2swift is available +if [ -f "packages/hf2swift/dist/cli.js" ]; then + for entry in "${GENERATED_MODELS[@]}"; do + model="${entry%%:*}" + output="${entry##*:}" + + # Regenerate model (swiftformat is already called by the generator) + node packages/hf2swift/dist/cli.js --model "$model" --output "$MODELS_DIR/$output" 2>/dev/null + done + + # Ensure consistent formatting with swiftformat (same as lint-staged) + if command -v swiftformat &> /dev/null; then + swiftformat "$MODELS_DIR" --quiet 2>/dev/null || true + fi + + # Check if any generated files changed + if ! git diff --quiet "$MODELS_DIR"/*Generated.swift 2>/dev/null; then + echo "❌ Generated Swift models are out of sync!" + echo "" + echo "The following generated files have uncommitted changes:" + git diff --name-only "$MODELS_DIR"/*Generated.swift + echo "" + echo "Either commit the regenerated files, or update the hf2swift generator" + echo "and regenerate with: pnpm hf2swift --model --output " + exit 1 + fi +fi + # Swift format check (if swift files changed) -if git diff --cached --name-only origin/main | grep -q '\.swift$'; then +if git diff --cached --name-only origin/main 2>/dev/null | grep -q '\.swift$'; then echo "→ SwiftFormat check..." cd packages/swift if command -v swiftformat &> /dev/null; then diff --git a/.prettierignore b/.prettierignore index 937e3ba..51b7aa5 100644 --- a/.prettierignore +++ b/.prettierignore @@ -1,5 +1,10 @@ dist/ -packages/swift/.build/ +build/ node_modules/ -pnpm-lock.yaml coverage/ + +packages/swift/.build/ +packages/docs-website/.source/ +pnpm-lock.yaml + +.venv/ diff --git a/README.md b/README.md index 29a6c74..a5ff47b 100644 --- a/README.md +++ b/README.md @@ -50,9 +50,9 @@ console.log(`${result.tokensPerSecond} tok/s`) | Provider | Models | Status | | --------- | ---------------- | ------------------ | | Qwen | Qwen3 0.6B–4B | ✅ **Recommended** | -| Microsoft | Phi-3.5, Phi-4 | ✅ High Quality | +| Microsoft | Phi-4 | ✅ High Quality | | Google | Gemma 3 1B–27B | ✅ Latest | -| Meta | Llama 3.2 | ✅ Auth required | +| Meta | Llama 4 | ✅ Auth required | | Mistral | Ministral 3B–14B | ✅ | | OpenAI | GPT-OSS 20B/120B | ✅ MoE | diff --git a/docs/rfcs/001-gpt-oss-moe-support.md b/docs/rfcs/001-gpt-oss-moe-support.md deleted file mode 100644 index d332a36..0000000 --- a/docs/rfcs/001-gpt-oss-moe-support.md +++ /dev/null @@ -1,260 +0,0 @@ -# RFC 001: GPT-OSS Mixture of Experts Support - -**Status**: Implemented -**Created**: 2026-01-09 -**Author**: node-mlx team - -## Summary - -Add support for OpenAI's GPT-OSS models (gpt-oss-20b, gpt-oss-120b) which use a Mixture of Experts (MoE) architecture. - -## Motivation - -GPT-OSS is OpenAI's first open-weight model family (released August 2025) under Apache 2.0 license. The 20B model is particularly attractive for local inference as it can run on systems with 16GB RAM while providing strong performance. - -**Available MLX Models**: - -- `mlx-community/gpt-oss-20b-MXFP4-Q8` (630k+ downloads) -- `mlx-community/gpt-oss-120b-MXFP4-Q8` -- Various quantization levels (4-bit, 8-bit) - -## Architecture Overview - -### Model Configuration - -```json -{ - "model_type": "gpt_oss", - "architectures": ["GptOssForCausalLM"], - "hidden_size": 2880, - "intermediate_size": 2880, - "num_hidden_layers": 24, - "num_attention_heads": 64, - "num_key_value_heads": 8, - "head_dim": 64, - "num_local_experts": 32, - "num_experts_per_tok": 4, - "sliding_window": 128, - "attention_bias": true, - "layer_types": ["sliding_attention", "full_attention", ...] -} -``` - -### Key Components - -#### 1. SwitchGLU (Mixture of Experts Layer) - -The core MoE component that routes tokens to selected experts: - -```python -class SwitchGLU: - def __init__(self, input_dims, hidden_dims, num_experts, activation, bias): - # Creates num_experts independent expert networks - # Each expert is a GLU (Gated Linear Unit) - pass - - def __call__(self, x, indices): - # Routes input x to experts specified by indices - # Returns weighted combination of expert outputs - pass -``` - -**Swift Implementation Required**: - -- `SwitchGLU` module with expert routing -- Batched expert computation for efficiency -- Weight loading for `experts.gate_proj`, `experts.up_proj`, `experts.down_proj` - -#### 2. Custom SwiGLU Activation - -GPT-OSS uses a modified SwiGLU with specific parameters: - -```python -def swiglu(x_linear, x_glu, alpha=1.702, limit=7.0): - x_glu = clip(x_glu, max=limit) - x_linear = clip(x_linear, min=-limit, max=limit) - glu_scaled = alpha * x_glu - sig = sigmoid(glu_scaled) - out_glu = x_glu * sig - return out_glu * (x_linear + 1) # Note: +1 bias -``` - -#### 3. Expert Router - -```python -class Router: - def __init__(self, hidden_size, num_experts): - self.linear = Linear(hidden_size, num_experts, bias=True) - - def __call__(self, x): - logits = self.linear(x) - # Select top-k experts - values, indices = topk(logits, k=num_experts_per_tok) - weights = softmax(values) - return weights, indices -``` - -#### 4. Attention with Sinks - -```python -class Attention: - def __init__(self): - self.sinks = zeros((num_attention_heads,)) # Learnable attention sinks - - def __call__(self, x, mask, cache): - # Standard attention with sink tokens for long context - output = scaled_dot_product_attention(q, k, v, sinks=self.sinks) - return output -``` - -#### 5. Mixed Attention Pattern - -Alternating between sliding window and full attention: - -```python -layer_types = ["sliding_attention", "full_attention"] * (num_layers // 2) -``` - -## Implementation Plan - -### Phase 1: Core MoE Infrastructure - -1. **Add `SwitchGLU` module** to Swift - - Implement expert weight storage - - Implement batched expert forward pass - - Handle quantized expert weights - -2. **Add TopK operator** for expert selection - - MLX Swift binding for `argpartition` - - Extract top-k indices and values - -3. **Implement SwiGLU activation** - - Custom activation with α=1.702, limit=7.0 - - Clipping and bias handling - -### Phase 2: GPT-OSS Model - -4. **Add `GptOssConfiguration`** struct - - All MoE-specific fields - - Layer type patterns - -5. **Add `GptOssAttention`** with sinks - - Learnable sink parameters - - Sliding/full attention switching - -6. **Add `GptOssMLP`** (Router + Experts) - - Expert routing logic - - SwitchGLU forward pass - -7. **Add `GptOssModel`** wrapper - - Weight sanitization for fused projections - - Cache creation with mixed types - -### Phase 3: Generator Support - -8. **Update `hf2swift` generator** - - Add MoE feature flags - - Generate SwitchGLU components - - Handle expert weight patterns - -## Estimated Effort - -| Component | Complexity | Time Estimate | -| -------------------- | ---------- | ---------------- | -| SwitchGLU module | High | 4-6 hours | -| TopK operator | Medium | 1-2 hours | -| SwiGLU activation | Low | 1 hour | -| Configuration | Low | 1 hour | -| Attention with sinks | Medium | 2-3 hours | -| MLP with routing | High | 3-4 hours | -| Model wrapper | Medium | 2 hours | -| Generator updates | Medium | 2-3 hours | -| Testing & debugging | High | 4-6 hours | -| **Total** | | **~20-28 hours** | - -## Open Questions - -1. **Expert parallelism**: Should we support multi-GPU expert sharding? -2. **Memory optimization**: Expert caching strategies for large models? -3. **Quantization**: How to handle per-expert quantization parameters? - -## References - -- [OpenAI GPT-OSS Announcement](https://openai.com/index/introducing-gpt-oss) -- [mlx-lm gpt_oss.py](https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/models/gpt_oss.py) -- [mlx-lm switch_layers.py](https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/models/switch_layers.py) - -## Implementation Notes - -The GPT-OSS MoE support has been implemented in the following files: - -### Swift Components - -1. **`MoELayers.swift`** - Core MoE infrastructure (manual implementation): - - `gptOssSwiGLU()` - Custom SwiGLU activation (α=1.702, limit=7.0) - - `MoERouter` - Token-to-expert routing with top-k selection via `argPartition` - - `SwitchGLU` - Batched expert computation - - `MoEMLP` - Complete MoE MLP layer - -2. **`GptOssGenerated.swift`** - **AUTO-GENERATED** by hf2swift: - - `GptOSSConfiguration` - Model configuration with MoE fields - - `GptOSSAttention` - Attention with learnable sinks - - `GptOSSDecoderLayer` - Decoder layer with MoE MLP - - `GptOSSModel` - Top-level model wrapper - -3. **`LLMModel.swift`** - Model registry updates: - - Added `gptOss` architecture case - - Model factory integration - -### Generator Updates (hf2swift) - -1. **`features.ts`** - MoE feature flags: - - `hasMoE`, `numExperts`, `numExpertsPerTok` - - `hasAttentionSinks`, `useCustomSwiGLU` - -2. **`config.ts`** - MoE configuration fields: - - `numLocalExperts`, `numExpertsPerTok`, `layerTypes` - -3. **`mlp.ts`** - MoE MLP generation: - - `generateMoEMlp()` function for MoE MLP components - -4. **`attention.ts`** - Attention sinks support: - - `sinks` parameter declaration and initialization - -5. **`model.ts`** - MoE-specific handling: - - `newCache()` using `layerTypes` for cache creation - - `sanitize()` with MoE expert weight mapping - -### Regenerate Command - -```bash -pnpm hf2swift --model gpt_oss --output packages/swift/Sources/NodeMLXCore/Models/GptOssGenerated.swift -``` - -### Usage - -```typescript -import { loadModel, generate } from "node-mlx" - -const model = await loadModel("mlx-community/gpt-oss-20b-MXFP4-Q8") -const response = await generate(model, "Hello, world!") -``` - -## Appendix: Weight Structure - -``` -model.embed_tokens.weight -model.layers.0.self_attn.q_proj.{weight,bias} -model.layers.0.self_attn.k_proj.{weight,bias} -model.layers.0.self_attn.v_proj.{weight,bias} -model.layers.0.self_attn.o_proj.{weight,bias} -model.layers.0.self_attn.sinks -model.layers.0.mlp.router.{weight,bias} -model.layers.0.mlp.experts.gate_proj.{weight,bias} # [num_experts, hidden, intermediate] -model.layers.0.mlp.experts.up_proj.{weight,bias} -model.layers.0.mlp.experts.down_proj.{weight,bias} -model.layers.0.input_layernorm.weight -model.layers.0.post_attention_layernorm.weight -model.norm.weight -lm_head.weight -``` diff --git a/docs/rfcs/002-ministral-smollm-lfm2-support.md b/docs/rfcs/002-ministral-smollm-lfm2-support.md deleted file mode 100644 index 3993efa..0000000 --- a/docs/rfcs/002-ministral-smollm-lfm2-support.md +++ /dev/null @@ -1,377 +0,0 @@ -# RFC 002: Ministral 3, SmolLM 3 & LFM2 Support - -**Status**: Implemented (Ministral 3 & SmolLM 3) / Deferred (LFM2) -**Created**: 2026-01-10 -**Author**: node-mlx team - -## Summary - -Add support for three new model families: - -1. **Ministral 3** (Mistral AI) - Multimodal edge-optimized models -2. **SmolLM 3** (Hugging Face) - Compact multilingual reasoning model -3. **LFM2** (Liquid AI) - Hybrid SSM/Transformer architecture - -## Model Overview - -### 1. Ministral 3 (Mistral AI) - -| Variant | Parameters | Context | Features | -| --------------- | ---------- | ------- | -------------------- | -| Ministral 3 3B | 3.4B | 256k | Vision, Multilingual | -| Ministral 3 8B | 8B | 256k | Vision, Multilingual | -| Ministral 3 14B | 14B | 256k | Vision, Multilingual | - -**Architecture**: Mistral-based with sliding window attention -**License**: Apache 2.0 -**Variants**: Base, Instruct, Reasoning - -**Expected config.json**: - -```json -{ - "model_type": "mistral", - "architectures": ["MistralForCausalLM"], - "hidden_size": 2560, - "num_hidden_layers": 32, - "num_attention_heads": 32, - "num_key_value_heads": 8, - "sliding_window": 4096, - "vocab_size": 131072 -} -``` - -### 2. SmolLM 3 (Hugging Face) - -| Variant | Parameters | Context | Features | -| ---------- | ---------- | ------- | -------------------------------- | -| SmolLM3-3B | 3B | 128k | 6 Languages, Think/NoThink modes | - -**Architecture**: Llama-based (likely `llama` or `smollm` model_type) -**License**: Apache 2.0 -**Languages**: English, French, Spanish, German, Italian, Portuguese - -**Expected Features**: - -- Long context (128k tokens) -- Dual reasoning modes ("think" vs "no_think") -- Efficient edge deployment - -**Expected config.json**: - -```json -{ - "model_type": "llama", - "architectures": ["LlamaForCausalLM"], - "hidden_size": 3072, - "num_hidden_layers": 36, - "num_attention_heads": 24, - "num_key_value_heads": 8, - "rope_theta": 1000000, - "max_position_embeddings": 131072 -} -``` - -### 3. LFM2 (Liquid AI) - -| Variant | Parameters | Features | -| --------- | ---------- | ---------------------- | -| LFM2-350M | 350M | Hybrid SSM/Transformer | -| LFM2-700M | 700M | Hybrid SSM/Transformer | -| LFM2-1.2B | 1.2B | Hybrid SSM/Transformer | -| LFM2-2.6B | 2.6B | Hybrid SSM/Transformer | - -**Architecture**: **Hybrid State Space Model + Transformer** -**License**: TBD (likely proprietary or restricted) -**Languages**: English, Japanese + 8 more - -⚠️ **Critical**: LFM2 uses a fundamentally different architecture combining: - -- State Space Model (SSM) layers (similar to Mamba) -- Transformer attention layers -- Custom hybrid routing - -This is NOT a standard Transformer architecture and requires significant new implementation work. - -## Implementation Analysis - -### Ministral 3 - -**Status**: ✅ Should work with existing Mistral support - -The generator already recognizes `ministral` in `features.ts` (line 230): - -```typescript -if (lower.includes("mistral") || lower.includes("ministral")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useSlidingWindow: true - // ... - } -} -``` - -**Required Work**: - -1. Verify model loads correctly with existing `MistralGenerated.swift` -2. Test quantized variants from `mlx-community` -3. Add Vision support (separate VLM implementation) - -**Estimated Effort**: 2-4 hours (mostly testing) - -### SmolLM 3 - -**Status**: ⚠️ May need minor adjustments - -SmolLM 3 is likely Llama-based but may have custom features. - -**Required Work**: - -1. Download model and inspect `config.json` for `model_type` -2. If `model_type: "llama"` → Should work with existing `LlamaGenerated.swift` -3. If custom `model_type: "smollm"` → Add feature flags in generator -4. Verify long context (128k) works with existing RoPE scaling - -**Potential Additions**: - -```typescript -// features.ts -if (lower.includes("smollm")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useSlidingWindow: false, - defaultRopeTheta: 1000000, // Long context - hasQKNorms: false, - normsPerLayer: 2 - // SmolLM specific if needed - } -} -``` - -**Estimated Effort**: 4-8 hours - -### LFM2 - -**Status**: ❌ Requires major new implementation - -LFM2 uses a Hybrid architecture that is NOT supported by the current codebase: - -#### State Space Model (SSM) Components Needed - -1. **Mamba/SSM Core**: - - Selective state space mechanism - - Hardware-efficient recurrence - - Different computational pattern than attention - -2. **Hybrid Layer Types**: - - ```python - layer_types = ["ssm", "attention", "ssm", "attention", ...] - ``` - -3. **New Modules Required**: - - `SSMLayer` - State space computation - - `SelectiveSSM` - Input-dependent state selection - - `CausalConv1d` - Causal convolution for SSM - - Hybrid model wrapper - -#### Architecture Comparison - -| Component | Transformer | SSM (Mamba-style) | -| ----------- | ----------------- | -------------------- | -| Core Op | Attention (O(n²)) | Recurrence (O(n)) | -| Memory | KV Cache | Hidden State | -| Parallelism | Fully parallel | Sequential (or scan) | - -**Estimated Effort**: 40-60 hours (new architecture) - -## Implementation Plan - -### Phase 1: Ministral 3 (Low effort) - -1. Test existing Mistral support with Ministral 3 models -2. Verify quantized variants work -3. Document any config differences -4. (Optional) Add Vision encoder support - -**Timeline**: 1 day - -### Phase 2: SmolLM 3 (Medium effort) - -1. Download and analyze SmolLM 3 config -2. Add `smollm` feature detection if needed -3. Generate and test Swift model -4. Verify 128k context support - -**Timeline**: 2-3 days - -### Phase 3: LFM2 (High effort) - Optional/Deferred - -⚠️ **Recommendation**: Defer LFM2 until: - -- Architecture details are publicly documented -- mlx-lm adds official support -- Community demand justifies the effort - -If proceeding: - -1. Research LFM2/Liquid architecture in detail -2. Implement SSM core modules -3. Add hybrid layer support to generator -4. Extensive testing and optimization - -**Timeline**: 2-3 weeks - -## Estimated Total Effort - -| Model | Complexity | Time | Priority | -| --------------------- | ---------- | ------------- | ----------- | -| Ministral 3 | Low | 2-4 hours | High | -| SmolLM 3 | Medium | 4-8 hours | High | -| LFM2 | Very High | 40-60 hours | Low (defer) | -| **Total (Phase 1+2)** | | **~1-2 days** | | - -## Open Questions - -1. **SmolLM 3 model_type**: Is it `llama`, `smollm`, or something else? -2. **Ministral 3 Vision**: Should we add multimodal support in this RFC? -3. **LFM2 Availability**: Are weights publicly available? What license? -4. **SSM Priority**: Is there community demand for Mamba/SSM support? - -## Recommendations - -1. **Proceed immediately** with Ministral 3 and SmolLM 3 -2. **Defer LFM2** until: - - mlx-lm adds official support (follow their implementation) - - Public weights and documentation available - - Clear demand from users - -3. **Consider separate RFC** for SSM/Mamba architecture support if LFM2 becomes priority - -## References - -- [Ministral 3 Collection](https://huggingface.co/collections/mistralai/ministral-3) -- [SmolLM 3 Repository](https://github.com/huggingface/smollm) -- [SmolLM 3 Website](https://smollm3.com/) -- [Liquid AI LFM2 Blog](https://www.liquid.ai/blog/introducing-lfm2-2-6b-redefining-efficiency-in-language-models) -- [Mamba Paper](https://arxiv.org/abs/2312.00752) (for SSM architecture reference) - -## Appendix: Quick Verification Commands - -### Test Ministral 3 - -```bash -# Check if existing Mistral support works -pnpm hf2swift --model mistral --output test-ministral.swift -cd packages/swift && swift build -``` - -### Inspect SmolLM 3 - -```bash -# Download and check config -huggingface-cli download HuggingFaceTB/SmolLM3-3B-Instruct config.json --local-dir ./tmp -cat ./tmp/config.json | jq '.model_type' -``` - -## Implementation Notes - -### Ministral 3 (Mistral 3) - -Implemented via generator with the following key features: - -**config.json Analysis**: - -```json -{ - "model_type": "mistral3", - "text_config": { - "model_type": "ministral3", - "hidden_size": 4096, - "num_hidden_layers": 34, - "num_attention_heads": 32, - "num_key_value_heads": 8, - "head_dim": 128, - "rope_theta": 1000000.0, - "rope_parameters": { - "rope_type": "yarn", - "factor": 16.0, - "mscale": 1.0, - "original_max_position_embeddings": 16384 - } - } -} -``` - -**Generator Features** (`features.ts`): - -- `hasYarnRope: true` - YaRN RoPE scaling for long context -- `defaultRopeTheta: 1000000` - 1M theta -- Standard Mistral-style attention and MLP - -**Generated Files**: - -- `Mistral3Generated.swift` - Full model implementation -- `RoPEParameters` struct for YaRN configuration - -### SmolLM 3 - -Implemented via generator with the following key features: - -**config.json Analysis**: - -```json -{ - "model_type": "smollm3", - "hidden_size": 2048, - "num_hidden_layers": 36, - "num_attention_heads": 16, - "num_key_value_heads": 4, - "rope_theta": 5000000.0, - "tie_word_embeddings": true, - "no_rope_layers": [1, 1, 1, 0, 1, 1, 1, 0, ...] -} -``` - -**Unique Feature**: `no_rope_layers` - Some layers skip RoPE entirely (1 = skip, 0 = use) - -**Generator Features** (`features.ts`): - -- `hasNoRopeLayers: true` - Layer-specific RoPE skipping -- `defaultRopeTheta: 5000000` - 5M theta for long context -- `hasWeightTying: true` - Shared embed/lm_head weights - -**Generated Files**: - -- `SmolLM3Generated.swift` - Full model implementation -- `shouldSkipRope(layerIdx)` helper in config -- Conditional RoPE application in attention - -### Regeneration Commands - -```bash -# Regenerate Ministral 3 -pnpm hf2swift --model mistral3 --output packages/swift/Sources/NodeMLXCore/Models/Mistral3Generated.swift - -# Regenerate SmolLM3 -pnpm hf2swift --model smollm3 --output packages/swift/Sources/NodeMLXCore/Models/SmolLM3Generated.swift - -# Verify build -cd packages/swift && swift build -c release -``` - -### Usage - -```typescript -import { loadModel, generate } from "node-mlx" - -// Ministral 3 -const ministral = await loadModel("mlx-community/Ministral-3-8B-Instruct-2512") -const response1 = await generate(ministral, "Hello!") - -// SmolLM3 -const smollm = await loadModel("HuggingFaceTB/SmolLM3-3B") -const response2 = await generate(smollm, "Explain quantum computing") -``` diff --git a/docs/rfcs/003-documentation-website.md b/docs/rfcs/003-documentation-website.md deleted file mode 100644 index b33d3d7..0000000 --- a/docs/rfcs/003-documentation-website.md +++ /dev/null @@ -1,647 +0,0 @@ -# RFC 003: Documentation Website & Open Source Marketing - -**Status**: Draft -**Created**: 2026-01-10 -**Author**: node-mlx team - -## Summary - -Erstellen einer modernen Dokumentations-Website für node-mlx mit starkem Marketing-Fokus, visueller Kommunikation und exzellenten Beispielen. Die Website wird über GitHub Pages veröffentlicht und ergänzt eine schlanke README. - -## Motivation - -Ein erfolgreiches Open-Source-Projekt braucht mehr als guten Code – es braucht: - -1. **Erste Sekunden zählen**: Entwickler entscheiden in Sekunden, ob ein Projekt interessant ist -2. **Visuelle Identität**: Logos, Screenshots und Grafiken schaffen Vertrauen -3. **Klare Wertversprechen**: Was macht node-mlx besonders? -4. **Einfacher Einstieg**: Von 0 zu funktionierendem Code in unter 2 Minuten - -### Aktuelle Probleme - -- README enthält zu viele Details (390+ Zeilen) -- Keine visuelle Identität -- Performance-Vorteile sind versteckt in Tabellen -- Keine interaktiven Demos oder Screenshots -- API-Dokumentation nicht durchsuchbar - -## Vorgeschlagene Lösung - -### 1. Dokumentations-Website (GitHub Pages) - -#### Tech Stack: Fumadocs + React Router + Vite - -**Warum Fumadocs mit React Router?** - -- Modernes Docs-Framework mit offizieller React Router-Unterstützung -- Kein Next.js nötig → einfacher Static Build für GitHub Pages -- TypeDoc-Integration für API-Dokumentation -- Exzellente Suche (eingebaut) -- Dark/Light Mode -- MDX für interaktive Komponenten -- Tailwind CSS für einfaches Styling -- Vite für blitzschnelle Builds - -**GitHub Pages Kompatibilität:** - -- `HashRouter` für clientseitiges Routing (`/#/docs/...`) -- Statischer Output → direkt deploybar -- Keine Server-Funktionen nötig - -#### Content-Struktur - -``` -content/ -├── docs/ -│ ├── index.mdx # Getting Started -│ ├── installation.mdx -│ ├── models/ -│ │ ├── qwen.mdx -│ │ ├── phi.mdx -│ │ ├── gemma.mdx -│ │ ├── llama.mdx -│ │ └── gpt-oss.mdx -│ ├── guides/ -│ │ ├── streaming.mdx -│ │ ├── memory-management.mdx -│ │ └── choosing-models.mdx -│ ├── api/ -│ │ └── [auto-generated by TypeDoc] -│ └── contributing.mdx -└── blog/ # Optional: Updates, Benchmarks - -public/ -├── logo.svg -├── logo-dark.svg -├── og-image.png # Social sharing (1200×630) -├── models/ # Provider logos -│ ├── qwen.svg -│ ├── phi.svg -│ ├── gemma.svg -│ └── llama.svg -├── screenshots/ -│ └── hero-terminal.png -└── icons/ - ├── apple-silicon.svg - ├── mlx.svg - └── nodejs.svg -``` - -### 2. Landing Page Design - -#### Hero Section - -``` -┌─────────────────────────────────────────────────────────────────┐ -│ │ -│ [Node.js Logo] × [MLX Logo] × [Apple Silicon Icon] │ -│ │ -│ ⚡ Run LLMs at Native Speed on Mac ⚡ │ -│ │ -│ The fastest way to run large language models in Node.js │ -│ Powered by Apple MLX. Built for Apple Silicon. │ -│ │ -│ ┌──────────────────────────────────────────────────────────┐ │ -│ │ $ npx node-mlx "What is 2+2?" │ │ -│ │ │ │ -│ │ ✓ Downloading Qwen3-4B-Instruct... │ │ -│ │ ⚡ Generated 24 tokens at 142 tok/s │ │ -│ │ │ │ -│ │ The answer is 4. │ │ -│ └──────────────────────────────────────────────────────────┘ │ -│ │ -│ [Get Started] [View on GitHub] [npm install node-mlx] │ -│ │ -└─────────────────────────────────────────────────────────────────┘ -``` - -#### Benefits Section - -``` -┌─────────────────────────────────────────────────────────────────┐ -│ Why node-mlx? │ -├─────────────────────┬─────────────────────┬────────────────────┤ -│ │ │ │ -│ 🚀 2× Faster │ 🧠 Unified Memory │ 📦 Zero Config │ -│ │ │ │ -│ vs node-llama-cpp │ No GPU memory │ npm install and │ -│ on Apple Silicon │ copying overhead │ you're ready │ -│ │ │ │ -├─────────────────────┼─────────────────────┼────────────────────┤ -│ │ │ │ -│ 🎯 TypeScript │ 🤗 HuggingFace │ 🔋 Efficient │ -│ │ │ │ -│ Full type safety │ Auto-download │ 4-bit quant │ -│ & IntelliSense │ from Hub │ native support │ -│ │ │ │ -└─────────────────────┴─────────────────────┴────────────────────┘ -``` - -#### Performance Visualization - -**Interaktiver Benchmark-Chart:** - -- Bar Chart: node-mlx vs node-llama-cpp -- Modelle: Mistral 7B, Phi-4 14B, Qwen3 4B, Gemma-3 12B -- Animierte Bars beim Scrollen -- Tooltip mit Details - -#### Model Showcase - -``` -┌─────────────────────────────────────────────────────────────────┐ -│ Supported Models │ -├─────────────────────────────────────────────────────────────────┤ -│ │ -│ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │ -│ │ [Qwen] │ │ [Phi] │ │ [Gemma] │ │ -│ │ Alibaba │ │ Microsoft │ │ Google │ │ -│ │ │ │ │ │ │ │ -│ │ 0.6B-4B │ │ 3.5-4 │ │ 1B-27B │ │ -│ │ ★ Default │ │ High qual. │ │ Latest │ │ -│ └─────────────┘ └─────────────┘ └─────────────┘ │ -│ │ -│ ┌─────────────┐ ┌─────────────┐ ┌─────────────┐ │ -│ │ [Llama] │ │ [Mistral] │ │ [GPT-OSS] │ │ -│ │ Meta │ │ Mistral │ │ OpenAI │ │ -│ │ │ │ │ │ │ │ -│ │ 1B-3B │ │ 3B-14B │ │ 20B-120B │ │ -│ │ Auth req. │ │ Ministral │ │ MoE │ │ -│ └─────────────┘ └─────────────┘ └─────────────┘ │ -│ │ -└─────────────────────────────────────────────────────────────────┘ -``` - -### 3. Visuelle Assets - -#### Benötigte Grafiken - -| Asset | Beschreibung | Format | -| --------------------- | ---------------------------------------- | ------ | -| `logo.svg` | node-mlx Logo (kombiniert N, MLX, Apple) | SVG | -| `logo-dark.svg` | Logo für Dark Mode | SVG | -| `og-image.png` | Social Media Preview (1200×630) | PNG | -| `favicon.ico` | Browser Tab Icon | ICO | -| `architecture.svg` | Vereinfachtes Architekturdiagramm | SVG | -| `benchmark-chart.svg` | Performance-Vergleich | SVG | -| `unified-memory.svg` | CPU/GPU unified memory Illustration | SVG | - -#### Model Provider Logos (Fair Use) - -- Qwen (Alibaba) -- Phi (Microsoft) -- Gemma (Google) -- Llama (Meta) -- Mistral -- GPT-OSS (OpenAI) - -**Hinweis**: Mit Attribution und Link zum Original - -### 4. Schlanke README - -Die README wird auf die Kernpunkte reduziert: - -````markdown -# node-mlx - -**The fastest way to run LLMs in Node.js on Apple Silicon.** - -[![CI](badge)][ci] [![npm](badge)][npm] [![License](badge)][license] - -## Quick Start - -\```bash -npm install node-mlx -npx node-mlx "Hello, world!" -\``` - -## Why node-mlx? - -🚀 **2× faster** than node-llama-cpp on Apple Silicon -🧠 **Unified memory** - no CPU/GPU copying -📦 **Zero config** - just npm install - -## Documentation - -📚 **[Full Documentation](https://sebastian-software.github.io/node-mlx/)** - -- [Getting Started](link) -- [Model Guide](link) -- [API Reference](link) -- [Performance Benchmarks](link) - -## Supported Models - -| Provider | Models | Status | -| --------- | ---------------- | ---------------- | -| Qwen | Qwen3 0.6B-4B | ✅ Default | -| Microsoft | Phi-3.5, Phi-4 | ✅ | -| Google | Gemma 3 1B-27B | ✅ | -| Meta | Llama 3.2 | ✅ Auth required | -| OpenAI | GPT-OSS 20B/120B | ✅ MoE | - -## Contributing - -See [CONTRIBUTING.md](./CONTRIBUTING.md) - -## License - -MIT © 2026 [Sebastian Software GmbH](https://sebastian-software.de) -```` - -### 5. GitHub Pages Deployment - -#### Workflow: `.github/workflows/docs.yml` - -```yaml -name: Deploy Docs - -on: - push: - branches: [main] - paths: - - "docs-website/**" - - "packages/node-mlx/src/**" # API changes trigger rebuild - workflow_dispatch: - -jobs: - build-and-deploy: - runs-on: ubuntu-latest - permissions: - contents: read - pages: write - id-token: write - - steps: - - uses: actions/checkout@v6 - - - uses: pnpm/action-setup@v4 - - - uses: actions/setup-node@v6 - with: - node-version: "22" - cache: "pnpm" - - - name: Install dependencies - run: pnpm install - - - name: Generate API docs (TypeDoc) - run: pnpm --filter docs-website typedoc - - - name: Build website (Vite) - run: pnpm --filter docs-website build - env: - BASE_URL: /node-mlx/ # GitHub Pages repo path - - - name: Setup Pages - uses: actions/configure-pages@v5 - - - name: Upload artifact - uses: actions/upload-pages-artifact@v3 - with: - path: "./docs-website/dist" # Vite output directory - - - name: Deploy to GitHub Pages - uses: actions/deploy-pages@v4 -``` - -#### Projektstruktur (Vite + React Router) - -``` -docs-website/ -├── package.json -├── vite.config.ts -├── tsconfig.json -├── tailwind.config.ts -├── index.html # Vite entry point -├── src/ -│ ├── main.tsx # React entry mit HashRouter -│ ├── App.tsx # Root component -│ ├── routes.tsx # Route definitions -│ ├── source.ts # Fumadocs content source -│ ├── components/ -│ │ ├── Layout.tsx -│ │ ├── Hero.tsx -│ │ ├── BenchmarkChart.tsx -│ │ └── ModelCard.tsx -│ └── styles/ -│ └── globals.css -├── content/ -│ └── docs/ # MDX documentation -│ ├── index.mdx -│ ├── installation.mdx -│ └── models/ -│ └── qwen.mdx -└── public/ - ├── logo.svg - └── models/ - └── qwen.svg -``` - -### 6. Beispiel-Seiten (MDX) - -#### Getting Started (`docs/index.mdx`) - -````mdx ---- -title: Getting Started -description: Run your first LLM in under 2 minutes ---- - -import { Steps, Callout } from "fumadocs-ui/components" -import { CopyButton } from "../components/CopyButton" - -# Getting Started - -node-mlx requires **macOS 14+** on **Apple Silicon** (M1/M2/M3/M4) - - - -### Install the package - -```bash -npm install node-mlx -``` -```` - -### Generate your first response - -```typescript -import { generate } from "node-mlx" - -const result = generate("qwen", "Explain quantum computing:", { - maxTokens: 200, - temperature: 0.7 -}) - -console.log(result.text) -// Quantum computing uses quantum bits (qubits) that can exist -// in multiple states simultaneously... -``` - -### Try the CLI - -```bash -npx node-mlx "What is 2+2?" -npx node-mlx --model phi --interactive # Chat mode -``` - - - -## What's Next? - -- [Choose the right model](/docs/models) for your use case -- [Understand memory management](/docs/guides/memory) -- [Explore the full API](/docs/api) - -```` - -#### Model Page (`docs/models/qwen.mdx`) - -```mdx ---- -title: Qwen Models -description: Alibaba's Qwen family - the recommended default ---- - -import { ModelCard, BenchmarkTable } from '../../components'; - -# Qwen Models - - - -## Overview - -Qwen3 is the **recommended default** model family for node-mlx. It offers -the best balance of quality, speed, and memory usage. - -## Available Variants - -| Model | Parameters | Memory | Speed | Best For | -|-------|------------|--------|-------|----------| -| `qwen-3-0.6b` | 600M | ~1.2 GB | 180 tok/s | Embedded, edge | -| `qwen-3-1.7b` | 1.7B | ~3 GB | 150 tok/s | General tasks | -| `qwen` | 4B | ~5 GB | 120 tok/s | **Recommended** | - -## Usage - -```typescript -import { loadModel } from "node-mlx" - -// Default (4B) - best balance -const model = loadModel("qwen") - -// Smaller variants -loadModel("qwen-3-0.6b") // Fastest -loadModel("qwen-3-1.7b") // Good balance - -// Legacy Qwen 2.5 -loadModel("qwen-2.5") // 1.5B -loadModel("qwen-2.5-3b") // 3B -```` - -## Performance - - - -## Tips - -- Start with the default `qwen` for most use cases -- Use `qwen-3-0.6b` for real-time applications -- Qwen excels at multilingual tasks (Chinese, English, etc.) - -```` - -## Implementation Plan - -### Phase 1: Foundation (1-2 Tage) - -1. **Setup docs-website package** (Vite + React Router + Fumadocs) - ```bash - # Im Monorepo - mkdir packages/docs-website - cd packages/docs-website - pnpm init - pnpm add react react-dom react-router-dom fumadocs-core fumadocs-ui - pnpm add -D vite @vitejs/plugin-react tailwindcss typescript -```` - -2. **Vite + React Router Konfiguration** - - `vite.config.ts` mit base path für GitHub Pages - - `HashRouter` für GitHub Pages Kompatibilität - - Tailwind mit Fumadocs Presets - -3. **Design System** - - Color Palette (Apple-inspiriert, Dark/Light Mode) - - Typography (SF Pro oder ähnlich) - - Custom Components - -4. **Logo & Branding** - - Logo Design (Node.js × MLX × Apple Silicon) - - Favicon - - OG Image für Social Sharing - -### Phase 2: Content (2-3 Tage) - -4. **Landing Page** - - Hero Section mit animated Terminal - - Benefits Grid - - Performance Chart (interaktiv) - - Model Showcase - -5. **Core Documentation** - - Getting Started - - Installation - - Model Guide (je Modell eine Seite) - - API Reference (TypeDoc) - -6. **Visual Assets** - - Model Provider Logos - - Architecture Diagram - - Benchmark Charts - -### Phase 3: Polish & Deploy (1 Tag) - -7. **GitHub Pages Workflow** - - docs.yml Action - - Auto-deploy on push - -8. **README Refactoring** - - Kürzen auf ~100 Zeilen - - Links zur Dokumentation - -9. **SEO & Meta** - - Open Graph tags - - sitemap.xml - - robots.txt - -## Estimated Effort - -| Phase | Aufgabe | Zeit | -| --------- | ------------------- | -------- | -| 1 | Fumadocs Setup | 4h | -| 1 | Design System | 4h | -| 1 | Logo & Branding | 2h | -| 2 | Landing Page | 6h | -| 2 | Documentation Pages | 8h | -| 2 | Visual Assets | 4h | -| 3 | GitHub Pages | 2h | -| 3 | README Refactor | 1h | -| 3 | SEO & Polish | 2h | -| **Total** | | **~33h** | - -## Open Questions - -1. ~~**Domain**: `docs.node-mlx.dev` vs `sebastian-software.github.io/node-mlx`?~~ - → **Entschieden: GitHub Pages** (`sebastian-software.github.io/node-mlx`) - -2. **Blog**: Sollen wir einen Blog für Updates integrieren? - → Nice-to-have, kann später ergänzt werden - -3. **Internationalisierung**: Deutsch + Englisch? - → Erstmal nur Englisch (internationale Zielgruppe) - -4. **Interactive Playground**: WASM-basierte Live-Demo möglich? - → Nicht für MLX (Apple Silicon only), aber Code-Snippets mit Copy-Button - -## Alternatives Considered - -### Fumadocs + Next.js (Original) - -**Pros:** - -- Mehr Features (ISR, API Routes) -- Größere Community - -**Cons:** - -- Server-Funktionen auf GitHub Pages nicht nutzbar -- Komplexerer Static Export -- Overhead für reine Dokumentation - -**→ Entschieden: React Router Version ist besser für GitHub Pages** - -### VitePress - -**Pros:** - -- Sehr schneller Build -- Einfach -- Gute Markdown-Erweiterungen - -**Cons:** - -- Vue-basiert (nicht React) -- Keine native TypeDoc-Integration -- Weniger Customization für Landing Page - -### Docusaurus - -**Pros:** - -- Sehr etabliert -- Viele Plugins -- React-basiert - -**Cons:** - -- Schwerer/langsamer -- Weniger modernes Design -- Mehr Boilerplate - -### Starlight (Astro) - -**Pros:** - -- Sehr performant -- Moderne DX -- Framework-agnostisch - -**Cons:** - -- Weniger TypeDoc-Support -- Kleinere Community -- Neue Syntax zu lernen - -## References - -- [Fumadocs](https://fumadocs.vercel.app/) -- [TypeDoc](https://typedoc.org/) -- [MLX Documentation](https://ml-explore.github.io/mlx/) -- [node-llama-cpp Docs](https://withcatai.github.io/node-llama-cpp/) -- [Vercel AI SDK Docs](https://sdk.vercel.ai/docs) (Design Inspiration) - -## Appendix: Design Inspiration - -### Color Palette - -```css -/* Light Mode */ ---brand-primary: #0066ff; /* Electric Blue */ ---brand-secondary: #ff6b35; /* Coral accent */ ---bg-primary: #fafafa; ---text-primary: #1a1a1a; - -/* Dark Mode */ ---brand-primary: #4d9fff; ---brand-secondary: #ff8c5a; ---bg-primary: #0d0d0d; ---text-primary: #f5f5f5; - -/* Apple-inspired gradients */ ---gradient-hero: linear-gradient(135deg, #667eea 0%, #764ba2 100%); ---gradient-metal: linear-gradient(180deg, #e8e8e8 0%, #c4c4c4 100%); -``` - -### Typography - -```css ---font-heading: "SF Pro Display", system-ui, sans-serif; ---font-body: "SF Pro Text", system-ui, sans-serif; ---font-mono: "SF Mono", "JetBrains Mono", monospace; -``` diff --git a/packages/benchmarks/src/full-benchmark.ts b/packages/benchmarks/src/full-benchmark.ts index c54b80a..d0b6876 100644 --- a/packages/benchmarks/src/full-benchmark.ts +++ b/packages/benchmarks/src/full-benchmark.ts @@ -65,14 +65,7 @@ const MODELS: Array<{ gguf: ".models/Qwen3-4B-Instruct-Q4_K_M.gguf" }, - // Phi (Microsoft) - { - family: "Phi", - name: "Phi-3.5 Mini", - size: "3.8B", - mlx: "mlx-community/Phi-3.5-mini-instruct-4bit", - gguf: ".models/Phi-3.5-mini-instruct-Q4_K_M.gguf" - }, + // Phi 4 (Microsoft) { family: "Phi", name: "Phi-4", diff --git a/packages/benchmarks/src/mlx-benchmark.ts b/packages/benchmarks/src/mlx-benchmark.ts index 3d1f648..a4c540f 100644 --- a/packages/benchmarks/src/mlx-benchmark.ts +++ b/packages/benchmarks/src/mlx-benchmark.ts @@ -27,13 +27,7 @@ const MODELS = [ id: "lmstudio-community/Qwen3-4B-Instruct-2507-MLX-4bit" }, - // Phi - { - family: "Phi", - name: "Phi-3.5 Mini", - size: "3.8B", - id: "mlx-community/Phi-3.5-mini-instruct-4bit" - }, + // Phi 4 { family: "Phi", name: "Phi-4", size: "14B", id: "mlx-community/phi-4-4bit" }, // Gemma 3 diff --git a/packages/benchmarks/src/mlx-models.ts b/packages/benchmarks/src/mlx-models.ts index 9ee0d43..a684a0f 100644 --- a/packages/benchmarks/src/mlx-models.ts +++ b/packages/benchmarks/src/mlx-models.ts @@ -22,8 +22,8 @@ interface Result { const MODELS = [ { id: "mlx-community/Qwen2.5-0.5B-Instruct-4bit", size: "0.5B" }, { id: "mlx-community/Qwen2.5-1.5B-Instruct-4bit", size: "1.5B" }, - { id: "mlx-community/Llama-3.2-1B-Instruct-4bit", size: "1B" }, - { id: "mlx-community/Phi-3-mini-4k-instruct-4bit", size: "3.8B" } + { id: "mlx-community/gemma-3-1b-it-4bit", size: "1B" }, + { id: "mlx-community/phi-4-4bit", size: "14B" } ] async function benchmark(modelId: string, size: string): Promise { diff --git a/packages/docs-website/app/routes/home.tsx b/packages/docs-website/app/routes/home.tsx index 255323a..cd25b93 100644 --- a/packages/docs-website/app/routes/home.tsx +++ b/packages/docs-website/app/routes/home.tsx @@ -52,10 +52,10 @@ const benchmarks = [ url: "https://qwenlm.github.io/blog/qwen3/" }, { - name: "Phi-3.5", - size: "3.8B", - nodemlx: 83, - llamacpp: 45, + name: "Phi-4", + size: "14B", + nodemlx: 45, + llamacpp: 24, logo: phiSvg, url: "https://azure.microsoft.com/en-us/products/phi" }, @@ -83,14 +83,6 @@ const benchmarks = [ logo: mistralSvg, url: "https://mistral.ai/technology/" }, - { - name: "Phi-4", - size: "14B", - nodemlx: 45, - llamacpp: 24, - logo: phiSvg, - url: "https://azure.microsoft.com/en-us/products/phi" - }, { name: "Gemma 3n", size: "2B", @@ -123,7 +115,7 @@ const models = [ name: "Phi", provider: "Microsoft", logo: phiSvg, - sizes: "3.5–4", + sizes: "14B", badge: "High Quality", url: "https://azure.microsoft.com/en-us/products/phi" }, @@ -139,7 +131,7 @@ const models = [ name: "Llama", provider: "Meta", logo: llamaSvg, - sizes: "1B–3B", + sizes: "17B–400B", badge: "Auth Required", url: "https://llama.meta.com/" }, diff --git a/packages/docs-website/content/docs/api/index.mdx b/packages/docs-website/content/docs/api/index.mdx index d0b3135..f57a0a9 100644 --- a/packages/docs-website/content/docs/api/index.mdx +++ b/packages/docs-website/content/docs/api/index.mdx @@ -121,13 +121,13 @@ interface Model { You can use short aliases or full HuggingFace paths: -| Alias | Full Path | -| --------- | --------------------------------------------- | -| `qwen` | `mlx-community/Qwen2.5-3B-Instruct-4bit` | -| `phi` | `mlx-community/Phi-4-mini-instruct-4bit` | -| `gemma` | `mlx-community/gemma-3-4b-it-4bit` | -| `llama` | `mlx-community/Llama-3.2-1B-Instruct-4bit` | -| `mistral` | `mlx-community/Mistral-7B-Instruct-v0.3-4bit` | +| Alias | Full Path | +| --------- | ---------------------------------------------------- | +| `qwen` | `lmstudio-community/Qwen3-4B-Instruct-2507-MLX-4bit` | +| `phi` | `mlx-community/phi-4-4bit` | +| `gemma` | `mlx-community/gemma-3-1b-it-4bit` | +| `llama` | `meta-llama/Llama-4-Scout-17B-16E-Instruct` | +| `mistral` | `mlx-community/Mistral-7B-Instruct-v0.3-4bit` | Or use any model from the [mlx-community](https://huggingface.co/mlx-community) on HuggingFace: diff --git a/packages/docs-website/content/docs/index.mdx b/packages/docs-website/content/docs/index.mdx index 9660253..be3671c 100644 --- a/packages/docs-website/content/docs/index.mdx +++ b/packages/docs-website/content/docs/index.mdx @@ -79,11 +79,11 @@ model.unload() ```typescript // First call - downloads and caches -const model = loadModel("mlx-community/Llama-3.2-1B-Instruct-4bit") +const model = loadModel("mlx-community/phi-4-4bit") // ⏳ Downloading... (one time only) // Second call - instant from cache -const model2 = loadModel("mlx-community/Llama-3.2-1B-Instruct-4bit") +const model2 = loadModel("mlx-community/phi-4-4bit") // ⚡ Ready immediately ``` diff --git a/packages/docs-website/content/docs/models/index.mdx b/packages/docs-website/content/docs/models/index.mdx index f45396f..aa0e806 100644 --- a/packages/docs-website/content/docs/models/index.mdx +++ b/packages/docs-website/content/docs/models/index.mdx @@ -25,14 +25,13 @@ loadModel("qwen-3-1.7b") // Small, good quality | `qwen-3-1.7b` | 1.7B | ~3 GB | 150 tok/s | General tasks | | `qwen` | 4B | ~5 GB | 120 tok/s | **Recommended** | -### Phi (Microsoft) +### Phi 4 (Microsoft) **Best for:** High quality reasoning, coding tasks ```typescript -loadModel("phi") // Phi-3.5-mini (default) -loadModel("phi3") // Phi-3-mini -loadModel("phi4") // Phi-4 (8GB, highest quality) +loadModel("phi") // Phi-4 (default, 14B, highest quality) +loadModel("phi4") // Phi-4 (alias) ``` ### Gemma 3 (Google) @@ -55,15 +54,15 @@ loadModel("gemma-3n") // Gemma-3n-E4B (default) loadModel("gemma-3n-e2b") // Gemma-3n-E2B (smaller) ``` -### Llama 3.2 (Meta) +### Llama 4 (Meta) -**Best for:** Well-tested, broad capability +**Best for:** Advanced reasoning, multilingual, large context > **Note:** Requires HuggingFace authentication. Run `huggingface-cli login` first. ```typescript -loadModel("llama") // Llama-3.2-1B-Instruct -loadModel("llama-3.2-3b") // Llama-3.2-3B-Instruct +loadModel("llama") // Llama-4-Scout (default) +loadModel("llama-4-scout") // Llama-4-Scout-17B-16E ``` ### GPT-OSS (OpenAI) @@ -82,7 +81,7 @@ You can use any compatible model from [mlx-community](https://huggingface.co/mlx ```typescript loadModel("mlx-community/Mistral-7B-Instruct-v0.3-4bit") loadModel("mlx-community/gemma-3-4b-it-4bit") -loadModel("mlx-community/Phi-3.5-mini-instruct-4bit") +loadModel("mlx-community/phi-4-4bit") ``` ## Model Quantization @@ -107,23 +106,23 @@ Most models come in two variants: - Smaller models where memory isn't a concern ```typescript -// 4-bit: ~2 GB RAM -loadModel("mlx-community/Llama-3.2-3B-Instruct-4bit") +// 4-bit: ~4 GB RAM +loadModel("mlx-community/phi-4-4bit") -// bf16: ~6 GB RAM -loadModel("mlx-community/Llama-3.2-3B-Instruct-bf16") +// bf16: ~28 GB RAM +loadModel("mlx-community/phi-4-bf16") ``` ## Supported Architectures -| Architecture | Example Models | Status | -| ------------ | --------------------- | --------------- | -| **Qwen2** | Qwen 2.5 | ✅ Full support | -| **Qwen3** | Qwen3 0.6B–4B | ✅ Full support | -| **Llama** | Llama 3.2, Mistral | ✅ Full support | -| **Phi3** | Phi-3, Phi-3.5, Phi-4 | ✅ Full support | -| **Gemma3** | Gemma 3 (1B–27B) | ✅ Full support | -| **Gemma3n** | Gemma 3n E2B/E4B | ✅ Full support | -| **Mistral3** | Ministral 3 (3B–14B) | ✅ Full support | -| **SmolLM3** | SmolLM3 3B | ✅ Full support | -| **GPT-OSS** | GPT-OSS 20B/120B MoE | ✅ Full support | +| Architecture | Example Models | Status | +| ------------ | -------------------- | --------------- | +| **Qwen2** | Qwen 2.5 | ✅ Full support | +| **Qwen3** | Qwen3 0.6B–4B | ✅ Full support | +| **Llama** | Llama 4, Mistral | ✅ Full support | +| **Phi3** | Phi-4 | ✅ Full support | +| **Gemma3** | Gemma 3 (1B–27B) | ✅ Full support | +| **Gemma3n** | Gemma 3n E2B/E4B | ✅ Full support | +| **Mistral3** | Ministral 3 (3B–14B) | ✅ Full support | +| **SmolLM3** | SmolLM3 3B | ✅ Full support | +| **GPT-OSS** | GPT-OSS 20B/120B MoE | ✅ Full support | diff --git a/packages/hf2swift/README.md b/packages/hf2swift/README.md new file mode 100644 index 0000000..2d00c29 --- /dev/null +++ b/packages/hf2swift/README.md @@ -0,0 +1,194 @@ +# hf2swift + +Swift code generator for HuggingFace transformer models. + +## Overview + +`hf2swift` generates Swift model implementations from HuggingFace model patterns. It produces code that integrates with Apple's MLX framework for efficient inference on Apple Silicon. + +## Features + +- **Feature-based generation**: Uses architectural features, not model names +- **Shared components**: Generates typealiases to reduce code duplication +- **Type-safe configs**: Generates Decodable configuration structs +- **SwiftFormat integration**: Consistent code formatting + +## Installation + +```bash +pnpm install +pnpm build +``` + +## Usage + +### CLI + +```bash +# Generate a model +pnpm hf2swift --model llama --output ./LlamaGenerated.swift + +# From config.json +pnpm hf2swift --config ./config.json --model llama --output ./LlamaGenerated.swift +``` + +### Programmatic + +```typescript +import { SwiftGenerator } from "@node-mlx/hf2swift" + +const generator = new SwiftGenerator("llama") +const swiftCode = generator.generate([]) +``` + +## Architecture + +``` +src/ +├── generator/ +│ ├── model-defs/ # Model family definitions +│ │ ├── types.ts # Interfaces + defaults +│ │ ├── llama.ts # Llama family +│ │ ├── qwen.ts # Qwen2, Qwen3 +│ │ ├── gemma.ts # Gemma3, Gemma3n +│ │ ├── phi.ts # Phi3, Phi4 +│ │ ├── mistral.ts # Mistral, Mistral3 +│ │ ├── gpt-oss.ts # GPT-OSS MoE +│ │ └── smollm.ts # SmolLM3 +│ ├── components/ # Swift code generators +│ │ ├── attention.ts # Attention layer +│ │ ├── mlp.ts # MLP layer +│ │ ├── decoder-layer.ts # Decoder layer +│ │ ├── model.ts # Model wrapper +│ │ └── rms-norm.ts # RMSNorm +│ ├── features.ts # Feature merging +│ ├── helpers.ts # Utility generators +│ └── index.ts # Main generator class +├── config.ts # Config struct generator +├── naming.ts # Name conversion utilities +└── cli.ts # CLI entry point +``` + +## Model Definitions + +Each model family is defined in `model-defs/`: + +```typescript +// model-defs/llama.ts +export const llama: ModelDefinition = { + name: "Llama", + matches: (modelType) => modelType.toLowerCase().includes("llama"), + architectural: { + rmsNormStyle: "standard", + activation: "silu", + hasQKNorms: false, + normsPerLayer: 2 + }, + configDefaults: { + ropeTheta: 10000, + rmsNormEps: 1e-5 + } +} +``` + +## Supported Models + +| Model | Type | Features | +| -------- | ---------- | ----------------------- | +| Llama | `llama` | Standard transformer | +| Qwen2 | `qwen2` | Attention bias | +| Qwen3 | `qwen3` | Q/K norms, weight tying | +| Phi-3/4 | `phi3` | Fused QKV/gate_up | +| Gemma3 | `gemma3` | 4 norms, Gemma RMSNorm | +| Gemma3n | `gemma3n` | AltUp, Laurel, VLM | +| Mistral | `mistral` | Sliding window | +| Mistral3 | `mistral3` | YaRN RoPE | +| SmolLM3 | `smollm3` | No-RoPE layers | +| GPT-OSS | `gpt_oss` | MoE, attention sinks | + +## Feature Flags + +### Architectural Features + +| Feature | Effect | +| ----------------------- | ----------------------------------- | +| `rmsNormStyle: "gemma"` | Uses (1+weight) scaling | +| `activation: "silu"` | SiLU activation in MLP | +| `hasFusedQKV: true` | Single qkv_proj instead of separate | +| `hasMoE: true` | Mixture of Experts MLP | +| `hasAltUp: true` | Alternating Updates (Gemma3n) | +| `hasQKNorms: true` | Q/K normalization | + +### Config Values (from config.json) + +| Value | Source | +| --------------- | ------------------- | +| `ropeTheta` | `rope_theta` | +| `slidingWindow` | `sliding_window` | +| `numExperts` | `num_local_experts` | + +## Generated Output + +### Simple Models (Llama, Qwen2) + +~195 lines using shared components: + +```swift +// MARK: - Attention +typealias LlamaAttention = StandardAttention + +// MARK: - MLP +typealias LlamaMLP = StandardMLP + +// MARK: - Decoder Layer +typealias LlamaDecoderLayer = StandardDecoderLayer +``` + +### Complex Models (Gemma3n) + +~700+ lines with custom implementations for advanced features. + +## Adding a New Model + +1. **Create definition**: `model-defs/newmodel.ts` + + ```typescript + export const newModel: ModelDefinition = { + name: "NewModel", + matches: (t) => t.includes("newmodel"), + architectural: { ...DEFAULT_ARCHITECTURAL, ... }, + configDefaults: { ...DEFAULT_CONFIG, ... } + } + ``` + +2. **Register**: Add to `model-defs/index.ts`: + + ```typescript + import { newModel } from "./newmodel.js" + const MODEL_REGISTRY = [..., newModel] + ``` + +3. **Test**: + ```bash + pnpm hf2swift --model newmodel + ``` + +## Development + +```bash +# Build +pnpm build + +# Test +pnpm test + +# Lint +pnpm lint + +# Watch mode +pnpm dev +``` + +## License + +MIT diff --git a/packages/hf2swift/package.json b/packages/hf2swift/package.json index cc401d0..8aa69a0 100644 --- a/packages/hf2swift/package.json +++ b/packages/hf2swift/package.json @@ -23,6 +23,7 @@ "devDependencies": { "@types/node": "^22.10.0", "tsup": "^8.5.1", + "tsx": "^4.21.0", "typescript": "^5.9.3", "vitest": "^4.0.16" }, diff --git a/packages/hf2swift/src/config.ts b/packages/hf2swift/src/config.ts index 814578b..89e18f6 100644 --- a/packages/hf2swift/src/config.ts +++ b/packages/hf2swift/src/config.ts @@ -28,8 +28,8 @@ export function generateConfigFromJson( parts.push(generateRoPEParametersStruct()) } - // Main configuration struct - parts.push(`public struct ${className}: Decodable, Sendable {`) + // Main configuration struct with BaseModelConfiguration conformance + parts.push(`public struct ${className}: Decodable, Sendable, BaseModelConfiguration {`) parts.push(generatePropertyDeclarations(features)) parts.push(generateHelperMethods(features)) parts.push(generateCodingKeys(features)) @@ -185,6 +185,11 @@ function generateHelperMethods(features?: ModelFeatures): string { if (features?.hasPerLayerIntermediateSize) { lines.push(` +/// Default intermediate size (first layer) for BaseModelConfiguration conformance +public var intermediateSize: Int { +intermediateSizes.first ?? 16384 +} + /// Get intermediate size for a specific layer public func intermediateSize(forLayer idx: Int) -> Int { if idx < intermediateSizes.count { @@ -373,14 +378,15 @@ intermediateSizes = Array(repeating: 16384, count: numHiddenLayers) lines.push("intermediateSize = try decode(.intermediateSize)") } - const defaultTheta = features?.defaultRopeTheta ?? 10000 + const defaultTheta = features?.ropeTheta ?? 10000 const defaultAttnBias = features?.hasAttentionBias ?? false const defaultMlpBias = features?.hasMlpBias ?? false + const defaultRmsNormEps = features?.rmsNormEps ?? 1e-6 lines.push(` vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) -rmsNormEps = try decode(.rmsNormEps, default: 1e-6) +rmsNormEps = try decode(.rmsNormEps, default: ${String(defaultRmsNormEps)}) ropeTheta = try decode(.ropeTheta, default: ${String(defaultTheta)}.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: ${String(defaultAttnBias)}) @@ -398,7 +404,10 @@ numExpertsPerTok = try decode(.numExpertsPerTok, default: ${numExpertsPerTok}) } if (features?.useSlidingWindow) { - lines.push("slidingWindow = try decode(.slidingWindow, default: 512)") + const defaultSlidingWindow = features.slidingWindow ?? 512 + lines.push( + `slidingWindow = try decode(.slidingWindow, default: ${String(defaultSlidingWindow)})` + ) if (!features.hasAltUp) { lines.push("slidingWindowPattern = try decode(.slidingWindowPattern, default: 6)") } diff --git a/packages/hf2swift/src/generator/components/attention.ts b/packages/hf2swift/src/generator/components/attention.ts index 03e20c1..0a8a210 100644 --- a/packages/hf2swift/src/generator/components/attention.ts +++ b/packages/hf2swift/src/generator/components/attention.ts @@ -23,94 +23,68 @@ export function generateAttention( if (features.hasFusedQKV) { return generateFusedQKVAttention(modelName, configClass, features) } + + // Check if model can use shared StandardAttention (no special features) + if (canUseSharedStandardAttention(features)) { + return generateSharedStandardAttention(modelName, configClass) + } + return generateStandardAttention(modelName, configClass, features) } /** - * Generate attention with fused QKV projection (Phi3, Phi4 style) + * Check if a model can use the shared StandardAttention implementation. + * Returns false if model needs any special attention features. */ -function generateFusedQKVAttention( - modelName: string, - configClass: string, - features: ModelFeatures -): string { - const scaleExpr = +function canUseSharedStandardAttention(features: ModelFeatures): boolean { + // These features require custom attention implementation + /* eslint-disable @typescript-eslint/prefer-nullish-coalescing -- logical OR for booleans */ + const hasSpecialFeatures = + features.useSlidingWindow || + features.hasKVSharing || + features.hasMoE || + features.hasNoRopeLayers || + features.hasQKNorms || + features.hasVNorm || + features.hasAttentionSinks || features.attentionScale !== undefined - ? String(features.attentionScale) - : "1.0 / sqrt(Float(headDim))" + /* eslint-enable @typescript-eslint/prefer-nullish-coalescing */ + + return !hasSpecialFeatures +} +/** + * Generate attention using shared StandardAttention. + * Used by Llama, Qwen2, and other simple models. + */ +function generateSharedStandardAttention(modelName: string, configClass: string): string { return ` // MARK: - Attention -class ${modelName}Attention: Module { -@ModuleInfo(key: "qkv_proj") var qkvProj: Linear -@ModuleInfo(key: "o_proj") var oProj: Linear - -let numHeads: Int -let numKVHeads: Int -let headDim: Int -let scale: Float -let rope: RoPE - -init(_ config: ${configClass}) { -self.numHeads = config.numAttentionHeads -self.numKVHeads = config.numKeyValueHeads -self.headDim = config.headDim -self.scale = ${scaleExpr} - -let qDim = numHeads * headDim -let kvDim = numKVHeads * headDim -let opSize = qDim + 2 * kvDim - -self._qkvProj.wrappedValue = Linear(config.hiddenSize, opSize, bias: false) -self._oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: false) -self.rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) +/// Standard attention - uses shared implementation +typealias ${modelName}Attention = StandardAttention<${configClass}> +` } -func callAsFunction( -_ hiddenStates: MLXArray, -mask: MLXFast.ScaledDotProductAttentionMaskMode, -cache: inout KVCache? -) -> MLXArray { -let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - -let qkv = qkvProj(hiddenStates) -let queryPos = numHeads * headDim -let kvPos = queryPos + numKVHeads * headDim - -var queries = qkv[0..., 0..., .. generic class. + * Only generates a typealias and protocol conformance. + */ +function generateFusedQKVAttention( + modelName: string, + configClass: string, + _features: ModelFeatures +): string { + return ` +// MARK: - Attention -// Update cache -if let c = cache { -(keys, values) = c.update(keys: keys, values: values) -} +/// AttentionConfiguration conformance for fused QKV attention +extension ${configClass}: AttentionConfiguration {} -// Attention using MLXFast (handles GQA automatically) -let output = MLXFast.scaledDotProductAttention( -queries: queries, -keys: keys, -values: values, -scale: scale, -mask: mask -) - -// Reshape back: [B, heads, L, headDim] -> [B, L, hidden] -let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) -return oProj(outputReshaped) -} -} +/// Fused QKV attention - uses shared implementation +typealias ${modelName}Attention = FusedQKVAttention<${configClass}> ` } @@ -246,18 +220,25 @@ function buildInitializations( } // RoPE initialization + const traditionalRope = features.useTraditionalRope ? "true" : "false" if (features.useSlidingWindow) { lines.push(`self.isSliding = !config.isGlobalLayer(layerIdx)`) const ropeBase = features.hasLocalRopeTheta ? "isSliding ? config.ropeLocalBaseFreq : config.ropeTheta" : "config.ropeTheta" lines.push(`let ropeBase = ${ropeBase}`) - lines.push(`self.rope = RoPE(dimensions: headDim, traditional: false, base: ropeBase)`) + lines.push( + `self.rope = RoPE(dimensions: headDim, traditional: ${traditionalRope}, base: ropeBase)` + ) } else if (features.hasNoRopeLayers) { lines.push(`self.skipRope = config.shouldSkipRope(layerIdx)`) - lines.push(`self.rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta)`) + lines.push( + `self.rope = RoPE(dimensions: headDim, traditional: ${traditionalRope}, base: config.ropeTheta)` + ) } else { - lines.push(`self.rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta)`) + lines.push( + `self.rope = RoPE(dimensions: headDim, traditional: ${traditionalRope}, base: config.ropeTheta)` + ) } // KV sharing @@ -300,7 +281,7 @@ function buildForwardBody(features: ModelFeatures): string { lines.push(`keys = kNorm(keys)`) } lines.push(`keys = keys.transposed(0, 2, 1, 3)`) - lines.push(`keys = rope.apply(keys, offset: offset)`) + lines.push(`keys = rope(keys, offset: offset)`) lines.push(`values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim])`) if (features.hasVNorm) { lines.push(`values = vNorm(values)`) @@ -310,7 +291,7 @@ function buildForwardBody(features: ModelFeatures): string { lines.push(`(keys, values) = c.update(keys: keys, values: values)`) lines.push(`}`) lines.push(`}`) - lines.push(`queries = rope.apply(queries, offset: offset)`) + lines.push(`queries = rope(queries, offset: offset)`) } else { // Standard path lines.push(`var keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim])`) @@ -334,12 +315,12 @@ function buildForwardBody(features: ModelFeatures): string { if (features.hasNoRopeLayers) { lines.push(`if !skipRope {`) - lines.push(`queries = rope.apply(queries, offset: offset)`) - lines.push(`keys = rope.apply(keys, offset: offset)`) + lines.push(`queries = rope(queries, offset: offset)`) + lines.push(`keys = rope(keys, offset: offset)`) lines.push(`}`) } else { - lines.push(`queries = rope.apply(queries, offset: offset)`) - lines.push(`keys = rope.apply(keys, offset: offset)`) + lines.push(`queries = rope(queries, offset: offset)`) + lines.push(`keys = rope(keys, offset: offset)`) } lines.push(``) diff --git a/packages/hf2swift/src/generator/components/decoder-layer.ts b/packages/hf2swift/src/generator/components/decoder-layer.ts index 9345a0b..b54574b 100644 --- a/packages/hf2swift/src/generator/components/decoder-layer.ts +++ b/packages/hf2swift/src/generator/components/decoder-layer.ts @@ -22,9 +22,55 @@ export function generateDecoderLayer( if (features.hasAltUp) { return generateAltUpDecoderLayer(modelName, configClass, features) } + + // Check if model can use shared StandardDecoderLayer + if (canUseSharedStandardDecoderLayer(features)) { + return generateSharedStandardDecoderLayer(modelName, configClass) + } + return generateStandardDecoderLayer(modelName, configClass, features) } +/** + * Check if a model can use the shared StandardDecoderLayer. + * Requires both StandardAttention and StandardMLP to be usable. + */ +function canUseSharedStandardDecoderLayer(features: ModelFeatures): boolean { + // Must be able to use both shared components + const canUseSharedAttention = + !features.useSlidingWindow && + !features.hasKVSharing && + !features.hasMoE && + !features.hasNoRopeLayers && + !features.hasQKNorms && + !features.hasVNorm && + !features.hasAttentionSinks && + features.attentionScale === undefined + + const canUseSharedMLP = + features.activation === "silu" && + !features.hasPerLayerIntermediateSize && + !features.hasSparseActivation && + !features.hasMoE + + // Also requires standard RMSNorm (not Gemma-style) + const usesStandardNorm = features.rmsNormStyle === "standard" && features.normsPerLayer === 2 + + return canUseSharedAttention && canUseSharedMLP && usesStandardNorm +} + +/** + * Generate decoder layer using shared StandardDecoderLayer. + */ +function generateSharedStandardDecoderLayer(modelName: string, configClass: string): string { + return ` +// MARK: - Decoder Layer + +/// Standard decoder layer - uses shared implementation +typealias ${modelName}DecoderLayer = StandardDecoderLayer<${configClass}> +` +} + function generateStandardDecoderLayer( modelName: string, configClass: string, diff --git a/packages/hf2swift/src/generator/components/mlp.ts b/packages/hf2swift/src/generator/components/mlp.ts index 3fc49d8..c70627c 100644 --- a/packages/hf2swift/src/generator/components/mlp.ts +++ b/packages/hf2swift/src/generator/components/mlp.ts @@ -33,6 +33,11 @@ export function generateMlp( return generateFusedGateUpMlp(modelName, configClass, activation) } + // Check if model can use shared StandardMLP (silu activation, no special features) + if (canUseSharedStandardMLP(features)) { + return generateSharedStandardMLP(modelName, configClass) + } + // eslint-disable-next-line @typescript-eslint/prefer-nullish-coalescing -- logical OR for booleans const needsLayerIdx = features.hasPerLayerIntermediateSize || features.hasSparseActivation const layerIdxParam = needsLayerIdx ? ", layerIdx: Int = 0" : "" @@ -96,6 +101,10 @@ return downProj(${activation}(gate) * up) ` } +/** + * Generate MLP with sparse gelu_topk activation. + * Uses shared MathUtils.erfinv instead of inline implementation. + */ function generateMlpWithSparseActivation( modelName: string, configClass: string, @@ -129,23 +138,12 @@ self.activationSparsity = 0.0 // Precompute std multiplier for gelu_topk if sparsity > 0 if activationSparsity > 0 { // sqrt(2) * erfinv(2 * sparsity - 1) -self.stdMultiplier = Float(sqrt(2.0)) * Self.erfinv(2.0 * activationSparsity - 1.0) +self.stdMultiplier = Float(sqrt(2.0)) * MathUtils.erfinv(2.0 * activationSparsity - 1.0) } else { self.stdMultiplier = nil } } -/// Approximate inverse error function -private static func erfinv(_ x: Float) -> Float { -let a: Float = 0.147 -let sign: Float = x < 0 ? -1 : 1 -let x2 = x * x -let lnTerm = log(1 - x2) -let term1 = 2 / (Float.pi * a) + lnTerm / 2 -let term2 = lnTerm / a -return sign * sqrt(sqrt(term1 * term1 - term2) - term1) -} - func callAsFunction(_ x: MLXArray) -> MLXArray { let gateOutput = gateProj(x) let activations: MLXArray @@ -167,59 +165,79 @@ return downProj(activations * upProj(x)) } function generateMoEMlp(modelName: string, configClass: string, features: ModelFeatures): string { - const useCustomSwiGLU = features.useCustomSwiGLU ?? false + // Determine which SwitchGLU variant to use + // GPT-OSS uses SwiGLU activation, others might use standard GELU + const expertsClass = features.useCustomSwiGLU ? "SwiGLUSwitchGLU" : "SwitchGLU" return ` // MARK: - MoE MLP -/// Mixture of Experts MLP using shared MoEMLP infrastructure +/// Mixture of Experts MLP with router and experts +/// Uses vendored SwitchLayers from mlx-swift-lm class ${modelName}MLP: Module { -@ModuleInfo(key: "router") var router: MoERouter -@ModuleInfo(key: "experts") var experts: SwitchGLU +@ModuleInfo(key: "experts") var experts: ${expertsClass} +@ModuleInfo(key: "router") var router: Linear -let numExperts: Int -let topK: Int +let hiddenSize: Int +let numLocalExperts: Int +let numExpertsPerTok: Int init(_ config: ${configClass}) { -self.numExperts = config.numLocalExperts -self.topK = config.numExpertsPerTok +hiddenSize = config.hiddenSize +numLocalExperts = config.numLocalExperts +numExpertsPerTok = config.numExpertsPerTok -_router.wrappedValue = MoERouter( -hiddenSize: config.hiddenSize, -numExperts: config.numLocalExperts, -topK: config.numExpertsPerTok, -bias: config.mlpBias -) -_experts.wrappedValue = SwitchGLU( +_experts.wrappedValue = ${expertsClass}( inputDims: config.hiddenSize, hiddenDims: config.intermediateSize, numExperts: config.numLocalExperts, -bias: config.mlpBias, -useCustomSwiGLU: ${String(useCustomSwiGLU)} +bias: config.mlpBias ) +_router.wrappedValue = Linear(config.hiddenSize, config.numLocalExperts, bias: config.mlpBias) } func callAsFunction(_ x: MLXArray) -> MLXArray { -let shape = x.shape -let batchSeq = shape.dropLast().reduce(1, *) -let hidden = shape.last! +let g = router(x) +let (expertScores, indices) = mlxTopK(g, k: numExpertsPerTok, axis: -1) +let expertWeights = softmax(expertScores, axis: -1, precise: true) -// Flatten to [batch * seq, hidden] -let xFlat = x.reshaped([batchSeq, hidden]) +var output = self.experts(x, indices: indices) -// Get routing weights and expert indices -let (weights, indices) = router(xFlat) +output = output * expandedDimensions(expertWeights, axis: -1) +return output.sum(axis: -2) +} +} +` +} -// Get expert outputs [batch * seq, topK, hidden] -let expertOutput = experts(xFlat, indices: indices) +/** + * Check if a model can use the shared StandardMLP implementation. + * StandardMLP uses SiLU activation and no special features. + */ +function canUseSharedStandardMLP(features: ModelFeatures): boolean { + // StandardMLP uses silu - skip for other activations + if (features.activation !== "silu") { + return false + } -// Weighted sum of expert outputs -let weightsExpanded = weights[.ellipsis, .newAxis] -let weightedOutput = sum(expertOutput * weightsExpanded, axis: 1) + // These features require custom MLP implementation + /* eslint-disable @typescript-eslint/prefer-nullish-coalescing -- logical OR for booleans */ + const hasSpecialFeatures = + features.hasPerLayerIntermediateSize || features.hasSparseActivation || features.hasMoE + /* eslint-enable @typescript-eslint/prefer-nullish-coalescing */ -// Reshape back to original shape -return weightedOutput.reshaped(shape) -} + return !hasSpecialFeatures } + +/** + * Generate MLP using shared StandardMLP. + * Used by Llama, Qwen2, and other simple models with SiLU activation. + */ +function generateSharedStandardMLP(modelName: string, configClass: string): string { + return ` +// MARK: - MLP + +/// Standard SwiGLU MLP - uses shared implementation +typealias ${modelName}MLP = StandardMLP<${configClass}> ` } diff --git a/packages/hf2swift/src/generator/components/model.ts b/packages/hf2swift/src/generator/components/model.ts index e67b5da..0a2d939 100644 --- a/packages/hf2swift/src/generator/components/model.ts +++ b/packages/hf2swift/src/generator/components/model.ts @@ -85,9 +85,10 @@ for (i, layerType) in layerTypes.prefix(cache.count).enumerated() { if layerType == "full_attention" { firstGlobalIdx = i; break } } let globalCache = firstGlobalIdx < cache.count ? cache[firstGlobalIdx] : nil -let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) -let firstSlidingCache = cache.first ?? nil -let slidingMask = createAttentionMask(h: hiddenStates, cache: firstSlidingCache, windowSize: slidingWindow)`, +let globalOffset = globalCache?.offset ?? 0 +let globalMask = createAttentionMask(n: hiddenStates.dim(1), offset: globalOffset, windowSize: nil) +let slidingOffset = cache.first??.offset ?? 0 +let slidingMask = createAttentionMask(n: hiddenStates.dim(1), offset: slidingOffset, windowSize: slidingWindow)`, layerLoop: `for i in 0.. 1 { -let firstCache = cache.first ?? nil -slidingMask = createAttentionMask(h: hiddenStates, cache: firstCache, windowSize: slidingWindow) +let slidingOffset = cache.first??.offset ?? 0 +slidingMask = createAttentionMask(n: hiddenStates.dim(1), offset: slidingOffset, windowSize: slidingWindow) } else { slidingMask = globalMask }`, @@ -126,7 +128,8 @@ self.slidingWindowPattern = config.slidingWindowPattern` } return { - maskHandling: `let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil)`, + maskHandling: `let offset = cache.first??.offset ?? 0 +let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil)`, layerLoop: `for i in 0.. [String: MLXArray] { -var result: [String: MLXArray] = [:] -for (key, value) in weights { -var newKey = key -if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } -else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } -else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } -if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } -${moeSanitization}result[newKey] = value +// Uses shared weight sanitization logic +return sanitizeWeights(weights) } -if result["lm_head.weight"] == nil { -for suffix in ["weight", "scales", "biases"] { -if let embedWeight = result["model.embed_tokens.\\(suffix)"] { result["lm_head.\\(suffix)"] = embedWeight } } +` } -return result + +function generateMoEModel(modelName: string, configClass: string, newCacheImpl: string): string { + return ` +// MARK: - Top-Level Model + +public class ${modelName}Model: Module, LLMModel { +public let vocabularySize: Int +public let numLayers: Int +public let numKVHeads: Int +public let headDim: Int +public let kvHeads: [Int] + +let model: ${modelName}ModelInner +private let configuration: ${configClass} +@ModuleInfo(key: "lm_head") var lmHead: Linear + +public var supportsCache: Bool { true } + +public init(_ config: ${configClass}) { +configuration = config +model = ${modelName}ModelInner(config) +vocabularySize = config.vocabSize +numLayers = config.numHiddenLayers +numKVHeads = config.numKeyValueHeads +headDim = config.headDim +kvHeads = (0 ..< config.numHiddenLayers).map { _ in config.numKeyValueHeads } +_lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) +} + +public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { +var cache: [KVCache?] = Array(repeating: nil, count: numLayers) +let hidden = model(inputIds, cache: &cache) +return lmHead(hidden) } + +public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { +var layerCaches: [KVCache?] +if let existingCache = cache { layerCaches = existingCache.map { $0 as KVCache? } } +else { layerCaches = Array(repeating: nil, count: numLayers) } +let hidden = model(inputIds, cache: &layerCaches) +cache = layerCaches.compactMap { $0 } +return lmHead(hidden) +} + +${newCacheImpl} + +${generateMoeSanitizeMethodInline()} } ` } +/** + * Generate MoE sanitize method that delegates to shared MoESanitizer. + * Reduces generated code from 80+ lines to 3 lines. + */ +function generateMoeSanitizeMethodInline(): string { + return `// MARK: - Weight Sanitization + +/// Sanitize MoE weights - delegates to shared implementation +public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { +MoESanitizer.sanitize(weights: weights) +}` +} + function buildNewCacheImpl(features: ModelFeatures): string { if (features.hasMoE) { return `public func newCache() -> [KVCache] { return (0.. MLXArray { -// Gemma uses (1 + weight) scaling -return MLXFast.rmsNorm(x, weight: 1 + weight, eps: eps) -} -} +/// Uses ported GemmaRMSNorm (1 + weight scaling) +typealias ${modelName}RMSNorm = GemmaRMSNorm `) } else { + // Standard models use the shared RMSNorm class parts.push(` -/// Standard RMSNorm -class ${modelName}RMSNorm: Module { -let eps: Float - -@ModuleInfo(key: "weight") var weight: MLXArray - -init(dimensions: Int, eps: Float = 1e-6) { -self.eps = eps -self._weight.wrappedValue = MLXArray.ones([dimensions]) -} - -func callAsFunction(_ x: MLXArray) -> MLXArray { -return MLXFast.rmsNorm(x, weight: weight, eps: eps) -} -} +/// Uses shared RMSNorm implementation +typealias ${modelName}RMSNorm = RMSNorm `) } diff --git a/packages/hf2swift/src/generator/features.test.ts b/packages/hf2swift/src/generator/features.test.ts index c54beb7..c4af618 100644 --- a/packages/hf2swift/src/generator/features.test.ts +++ b/packages/hf2swift/src/generator/features.test.ts @@ -2,66 +2,132 @@ import { describe, it, expect } from "vitest" import { getModelFeatures } from "./features.js" describe("getModelFeatures", () => { - it("returns Gemma3 features for gemma3 model", () => { - const features = getModelFeatures("gemma3") - - expect(features.rmsNormStyle).toBe("gemma") - expect(features.activation).toBe("geluApproximate") - expect(features.useClipResidual).toBe(true) - expect(features.useSlidingWindow).toBe(true) - expect(features.defaultRopeTheta).toBe(1000000) - expect(features.hasLocalRopeTheta).toBe(true) - expect(features.useEmbeddingScale).toBe(true) - expect(features.hasQKNorms).toBe(true) - expect(features.normsPerLayer).toBe(4) - }) + describe("architectural features (model-specific)", () => { + it("returns Gemma3 architectural features", () => { + const features = getModelFeatures("gemma3") - it("returns Gemma3 features for gemma-3 model", () => { - const features = getModelFeatures("gemma-3") + expect(features.rmsNormStyle).toBe("gemma") + expect(features.activation).toBe("geluApproximate") + expect(features.useClipResidual).toBe(true) + expect(features.useEmbeddingScale).toBe(true) + expect(features.hasQKNorms).toBe(true) + expect(features.normsPerLayer).toBe(4) + }) - expect(features.rmsNormStyle).toBe("gemma") - expect(features.activation).toBe("geluApproximate") - }) + it("returns Gemma3 features for gemma-3 variant", () => { + const features = getModelFeatures("gemma-3") - it("returns Qwen features for qwen2 model", () => { - const features = getModelFeatures("qwen2") + expect(features.rmsNormStyle).toBe("gemma") + expect(features.activation).toBe("geluApproximate") + }) - expect(features.rmsNormStyle).toBe("standard") - expect(features.activation).toBe("silu") - expect(features.useClipResidual).toBe(false) - expect(features.useSlidingWindow).toBe(false) - expect(features.normsPerLayer).toBe(2) - }) + it("returns Qwen2 architectural features", () => { + const features = getModelFeatures("qwen2") - it("returns Llama features for llama model", () => { - const features = getModelFeatures("llama") + expect(features.rmsNormStyle).toBe("standard") + expect(features.activation).toBe("silu") + expect(features.useClipResidual).toBe(false) + expect(features.normsPerLayer).toBe(2) + }) - expect(features.rmsNormStyle).toBe("standard") - expect(features.activation).toBe("silu") - expect(features.useSlidingWindow).toBe(false) - }) + it("returns Llama architectural features", () => { + const features = getModelFeatures("llama") + + expect(features.rmsNormStyle).toBe("standard") + expect(features.activation).toBe("silu") + }) + + it("returns Phi architectural features with fused projections", () => { + const features = getModelFeatures("phi3") - it("returns Phi features for phi model", () => { - const features = getModelFeatures("phi3") + expect(features.rmsNormStyle).toBe("standard") + expect(features.activation).toBe("silu") + expect(features.hasFusedQKV).toBe(true) + expect(features.hasFusedGateUp).toBe(true) + }) - expect(features.rmsNormStyle).toBe("standard") - expect(features.activation).toBe("silu") + it("returns GPT-OSS architectural features with MoE", () => { + const features = getModelFeatures("gpt_oss") + + expect(features.hasMoE).toBe(true) + expect(features.hasAttentionSinks).toBe(true) + expect(features.useCustomSwiGLU).toBe(true) + }) + + it("returns default features for unknown models", () => { + const features = getModelFeatures("unknown_model") + + expect(features.rmsNormStyle).toBe("standard") + expect(features.activation).toBe("gelu") + expect(features.useClipResidual).toBe(false) + }) }) - it("returns Mistral features with sliding window", () => { - const features = getModelFeatures("mistral") + describe("config values (from defaults)", () => { + it("returns Gemma3 default config values", () => { + const features = getModelFeatures("gemma3") + + expect(features.useSlidingWindow).toBe(true) + expect(features.ropeTheta).toBe(1000000) + expect(features.hasLocalRopeTheta).toBe(true) + }) + + it("returns Mistral default config values with sliding window", () => { + const features = getModelFeatures("mistral") + + expect(features.useSlidingWindow).toBe(true) + }) - expect(features.rmsNormStyle).toBe("standard") - expect(features.activation).toBe("silu") - expect(features.useSlidingWindow).toBe(true) + it("returns GPT-OSS default MoE config values", () => { + const features = getModelFeatures("gpt_oss") + + expect(features.numExperts).toBe(128) + expect(features.numExpertsPerTok).toBe(4) + expect(features.slidingWindow).toBe(128) + expect(features.ropeTheta).toBe(150000) + }) }) - it("returns default features for unknown models", () => { - const features = getModelFeatures("unknown_model") + describe("config override (from config.json)", () => { + it("overrides ropeTheta from config.json", () => { + const features = getModelFeatures("llama", { rope_theta: 500000 }) + + expect(features.ropeTheta).toBe(500000) + }) + + it("overrides attentionBias from config.json", () => { + const features = getModelFeatures("llama", { attention_bias: true }) + + expect(features.hasAttentionBias).toBe(true) + }) + + it("overrides mlpBias from config.json", () => { + const features = getModelFeatures("llama", { mlp_bias: true }) + + expect(features.hasMlpBias).toBe(true) + }) + + it("overrides slidingWindow from config.json", () => { + const features = getModelFeatures("llama", { sliding_window: 4096 }) + + expect(features.slidingWindow).toBe(4096) + expect(features.useSlidingWindow).toBe(true) + }) + + it("overrides numExperts from config.json", () => { + const features = getModelFeatures("gpt_oss", { num_local_experts: 64 }) + + expect(features.numExperts).toBe(64) + }) + + it("preserves architectural features when config provided", () => { + const features = getModelFeatures("gemma3", { rope_theta: 999999 }) - expect(features.rmsNormStyle).toBe("standard") - expect(features.activation).toBe("gelu") - expect(features.useClipResidual).toBe(false) - expect(features.useSlidingWindow).toBe(false) + // Config override + expect(features.ropeTheta).toBe(999999) + // Architectural preserved + expect(features.rmsNormStyle).toBe("gemma") + expect(features.activation).toBe("geluApproximate") + }) }) }) diff --git a/packages/hf2swift/src/generator/features.ts b/packages/hf2swift/src/generator/features.ts index 55c5d0f..f930567 100644 --- a/packages/hf2swift/src/generator/features.ts +++ b/packages/hf2swift/src/generator/features.ts @@ -1,331 +1,110 @@ /** * Model-specific feature flags for code generation * - * These flags control which Swift code patterns are generated - * for different model architectures. - */ - -/** - * Model-specific feature configuration + * Two-tier system: + * 1. Architectural features - determined by model family (from model-defs/) + * 2. Config values - read from config.json with model-specific defaults + * + * This separation ensures: + * - Model-specific code paths are feature-driven, not name-driven + * - Config values come from the source of truth (config.json) + * - Reasonable defaults when config values are missing */ -export interface ModelFeatures { - // === Core Architecture === - - /** RMSNorm style: "gemma" uses (1+weight), "standard" uses weight directly */ - rmsNormStyle: "gemma" | "standard" - - /** Activation function: "gelu", "geluApproximate" (Gemma), or "silu" */ - activation: "gelu" | "geluApproximate" | "silu" - - /** Use clipResidual for float16 overflow protection */ - useClipResidual: boolean - - /** Sliding window attention support */ - useSlidingWindow: boolean - - /** Default RoPE theta (10000 for most, 1000000 for Gemma3) */ - defaultRopeTheta: number - - /** Has separate local RoPE theta for sliding window layers */ - hasLocalRopeTheta: boolean - - /** Gemma-style embedding scaling (multiply by sqrt(hiddenSize)) */ - useEmbeddingScale: boolean - - /** Has Q/K norms before attention */ - hasQKNorms: boolean - - /** Number of norms per decoder layer (2 for most, 4 for Gemma3) */ - normsPerLayer: 2 | 4 - - /** Has attention bias (read from config.attention_bias, default varies by model) */ - hasAttentionBias?: boolean - - /** Has MLP bias (read from config.mlp_bias, default false) */ - hasMlpBias?: boolean - - // === Advanced Features (Gemma3n and future models) === - - /** AltUp (Alternating Updates) for efficient sparse computation */ - hasAltUp?: boolean - - /** Laurel (Learned Augmented Residual) blocks */ - hasLaurel?: boolean - - /** Per-layer input embeddings */ - hasPerLayerInputs?: boolean - - /** KV-cache sharing for later layers */ - hasKVSharing?: boolean - - /** Per-layer intermediate MLP sizes (array instead of single value) */ - hasPerLayerIntermediateSize?: boolean - /** Sparse activation with gelu_topk */ - hasSparseActivation?: boolean +// Re-export types from model-defs +export type { ArchitecturalFeatures, ConfigValues, ModelFeatures } from "./model-defs/index.js" - /** Value normalization (RMSNoScale) in attention */ - hasVNorm?: boolean +// Re-export isGemma3n for backward compatibility +export { isGemma3n } from "./model-defs/index.js" - /** Weight tying (use embed_tokens.weight for lm_head) */ - hasWeightTying?: boolean - - /** Logit softcapping */ - hasLogitSoftcapping?: boolean - - /** Attention scale override (e.g., 1.0 for Gemma3n instead of 1/sqrt(headDim)) */ - attentionScale?: number - - // === Fused Projections === - - /** Use fused QKV projection instead of separate q_proj, k_proj, v_proj */ - hasFusedQKV?: boolean - - /** Use fused gate_up_proj instead of separate gate_proj, up_proj */ - hasFusedGateUp?: boolean - - // === Mixture of Experts (MoE) === - - /** Uses Mixture of Experts architecture */ - hasMoE?: boolean - - /** Number of expert networks */ - numExperts?: number - - /** Number of experts selected per token */ - numExpertsPerTok?: number - - /** Has learnable attention sinks */ - hasAttentionSinks?: boolean - - /** Uses custom SwiGLU activation (alpha=1.702, limit=7.0) */ - useCustomSwiGLU?: boolean - - // === SmolLM3 / Ministral 3 specific === - - /** Some layers skip RoPE (SmolLM3 no_rope_layers config) */ - hasNoRopeLayers?: boolean - - /** Uses YaRN RoPE scaling (Ministral 3) */ - hasYarnRope?: boolean -} +import { + getArchitecturalFeatures, + getDefaultConfigValues, + type ConfigValues, + type ModelFeatures +} from "./model-defs/index.js" /** - * Check if model is Gemma 3n + * Raw config.json structure (partial) */ -export function isGemma3n(modelType: string): boolean { - const lower = modelType.toLowerCase() - return lower.includes("gemma3n") || lower.includes("gemma-3n") || lower.includes("gemma_3n") +interface ConfigJson { + rope_theta?: number + attention_bias?: boolean + mlp_bias?: boolean + rms_norm_eps?: number + sliding_window?: number | null + num_local_experts?: number + num_experts_per_tok?: number + tie_word_embeddings?: boolean + rope_local_base_freq?: number } /** - * Get default features for a model type + * Extract config values from config.json + * Returns only values that are explicitly set */ -export function getModelFeatures(modelType: string): ModelFeatures { - const lower = modelType.toLowerCase() - - // Gemma 3n - Very specialized architecture (check first!) - if (isGemma3n(modelType)) { - return { - // Base features (similar to Gemma 3) - rmsNormStyle: "standard", // Gemma3n uses standard RMSNorm (not 1+weight) - activation: "geluApproximate", - useClipResidual: false, - useSlidingWindow: true, - defaultRopeTheta: 1000000, - hasLocalRopeTheta: true, - useEmbeddingScale: true, - hasQKNorms: true, - normsPerLayer: 4, +function extractConfigValues(configJson: ConfigJson): Partial { + const values: Partial = {} - // Gemma 3n specific advanced features - hasAltUp: true, - hasLaurel: true, - hasPerLayerInputs: true, - hasKVSharing: true, - hasPerLayerIntermediateSize: true, - hasSparseActivation: true, - hasVNorm: true, - hasWeightTying: true, - hasLogitSoftcapping: true, - attentionScale: 1.0 - } + if (configJson.rope_theta !== undefined) { + values.ropeTheta = configJson.rope_theta } - // Gemma 3 - Advanced features - if (lower.includes("gemma3") || lower.includes("gemma-3")) { - return { - rmsNormStyle: "gemma", - activation: "geluApproximate", - useClipResidual: true, - useSlidingWindow: true, - defaultRopeTheta: 1000000, - hasLocalRopeTheta: true, - useEmbeddingScale: true, - hasQKNorms: true, - normsPerLayer: 4 - } + if (configJson.attention_bias !== undefined) { + values.hasAttentionBias = configJson.attention_bias } - // Qwen3 - Like Qwen2 but with Q/K norms and no attention bias - if (lower.includes("qwen3")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 1000000, // Qwen3 uses 1M rope theta - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: true, // Qwen3 has Q/K norms - normsPerLayer: 2, - hasAttentionBias: false, // Qwen3 has no attention bias - hasMlpBias: false, - hasWeightTying: true // Qwen3 uses tie_word_embeddings - } + if (configJson.mlp_bias !== undefined) { + values.hasMlpBias = configJson.mlp_bias } - // Qwen2 - Standard with SiLU, has attention bias by default - if (lower.includes("qwen")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2, - hasAttentionBias: true, // Qwen2/2.5 has attention_bias: true by default - hasMlpBias: false - } + if (configJson.rms_norm_eps !== undefined) { + values.rmsNormEps = configJson.rms_norm_eps } - // Llama - Standard with SiLU - if (lower.includes("llama")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2 - } + if (configJson.sliding_window !== undefined && configJson.sliding_window !== null) { + values.slidingWindow = configJson.sliding_window + values.useSlidingWindow = true } - // Phi3/Phi4 - Fused projections and SiLU - if (lower.includes("phi")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2, - hasFusedQKV: true, // Phi3/Phi4 uses qkv_proj instead of separate q/k/v - hasFusedGateUp: true // Phi3/Phi4 uses gate_up_proj instead of separate gate/up - } + if (configJson.num_local_experts !== undefined) { + values.numExperts = configJson.num_local_experts } - // Mistral - with sliding window - if (lower.includes("mistral") || lower.includes("ministral")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: true, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2 - } + if (configJson.num_experts_per_tok !== undefined) { + values.numExpertsPerTok = configJson.num_experts_per_tok } - // GPT-OSS - Mixture of Experts with custom SwiGLU - if (lower.includes("gpt_oss") || lower.includes("gptoss") || lower.includes("gpt-oss")) { - return { - rmsNormStyle: "standard", - activation: "silu", // Uses custom SwiGLU but base is silu - useClipResidual: false, - useSlidingWindow: true, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2, - hasAttentionBias: true, - hasMlpBias: true, - // MoE features - hasMoE: true, - numExperts: 32, - numExpertsPerTok: 4, - hasAttentionSinks: true, - useCustomSwiGLU: true - } + if (configJson.tie_word_embeddings !== undefined) { + values.hasWeightTying = configJson.tie_word_embeddings } - // SmolLM3 - Compact multilingual model with no_rope_layers - if (lower.includes("smollm3") || lower.includes("smollm-3") || lower.includes("smollm_3")) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 5000000, // 5M theta - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2, - hasAttentionBias: false, - hasMlpBias: false, - hasWeightTying: true, - // SmolLM3 specific: some layers skip RoPE (handled via no_rope_layers config) - hasNoRopeLayers: true - } + if (configJson.rope_local_base_freq !== undefined) { + values.hasLocalRopeTheta = true } - // Mistral 3 / Ministral 3 - Multimodal with YaRN RoPE - if ( - lower.includes("mistral3") || - lower.includes("mistral-3") || - lower.includes("ministral3") || - lower.includes("ministral-3") - ) { - return { - rmsNormStyle: "standard", - activation: "silu", - useClipResidual: false, - useSlidingWindow: false, // Ministral 3 uses full attention - defaultRopeTheta: 1000000, // 1M theta - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2, - hasAttentionBias: false, - hasMlpBias: false, - // YaRN RoPE scaling - hasYarnRope: true - } - } + return values +} - // Default features +/** + * Get complete model features + * + * @param modelType - Model type name (e.g., "gemma3", "qwen2", "llama") + * @param configJson - Optional config.json contents to extract values from + * @returns Combined architectural features and config values + */ +export function getModelFeatures( + modelType: string, + configJson?: Record +): ModelFeatures { + const architectural = getArchitecturalFeatures(modelType) + const defaults = getDefaultConfigValues(modelType) + const fromConfig = configJson ? extractConfigValues(configJson as ConfigJson) : {} + + // Merge: architectural + defaults + config (config wins) return { - rmsNormStyle: "standard", - activation: "gelu", - useClipResidual: false, - useSlidingWindow: false, - defaultRopeTheta: 10000, - hasLocalRopeTheta: false, - useEmbeddingScale: false, - hasQKNorms: false, - normsPerLayer: 2 + ...architectural, + ...defaults, + ...fromConfig } } diff --git a/packages/hf2swift/src/generator/helpers.ts b/packages/hf2swift/src/generator/helpers.ts index 78ab974..c0167ef 100644 --- a/packages/hf2swift/src/generator/helpers.ts +++ b/packages/hf2swift/src/generator/helpers.ts @@ -25,17 +25,21 @@ import MLXNN` export function generateHelpers(features: ModelFeatures): string { const parts: string[] = ["// MARK: - Utility Functions"] - // clipResidual helper for float16 overflow protection + // clipResidual helper - use shared MathUtils if (features.useClipResidual) { parts.push(` -/// Clip residual for float16 overflow protection (matching mlx-lm) +/// Clip residual for float16 overflow protection - uses shared implementation private func clipResidual(_ x: MLXArray, _ y: MLXArray) -> MLXArray { - if x.dtype != .float16 { - return x + y - } - let bound = Float16.greatestFiniteMagnitude - let sum = (x.asType(.float32) + y.asType(.float32)) - return clip(sum, min: MLXArray(-Float(bound)), max: MLXArray(Float(bound))).asType(.float16) + MathUtils.clipResidual(x, y) +}`) + } + + // mlxTopK helper - use shared MathUtils + if (features.hasMoE) { + parts.push(` +/// Top-k selection for MoE routing - uses shared implementation +private func mlxTopK(_ a: MLXArray, k: Int, axis: Int = -1) -> (values: MLXArray, indices: MLXArray) { + MathUtils.topK(a, k: k, axis: axis) }`) } @@ -44,133 +48,30 @@ private func clipResidual(_ x: MLXArray, _ y: MLXArray) -> MLXArray { /** * Generate Laurel (Learned Augmented Residual) block - * Low-rank residual layer for efficient computation + * Uses shared LaurelBlock component. */ -export function generateLaurelBlock( - modelName: string, - configClass: string, - normType: string -): string { +export function generateLaurelBlock(modelName: string, configClass: string): string { return `// MARK: - Laurel Block -/// Low-rank residual layer (Learned Augmented Residual) -/// Note: This layer adds the residual internally (returns x + laurel_output) -class ${modelName}LaurelBlock: Module { - @ModuleInfo(key: "linear_left") var linearLeft: Linear - @ModuleInfo(key: "linear_right") var linearRight: Linear - @ModuleInfo(key: "post_laurel_norm") var postLaurelNorm: ${normType} - - init(_ config: ${configClass}) { - _linearLeft.wrappedValue = Linear(config.hiddenSize, config.laurelRank, bias: false) - _linearRight.wrappedValue = Linear(config.laurelRank, config.hiddenSize, bias: false) - _postLaurelNorm.wrappedValue = ${normType}(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } +/// LaurelConfiguration conformance for shared LaurelBlock +extension ${configClass}: LaurelConfiguration {} - func callAsFunction(_ x: MLXArray) -> MLXArray { - var laurel = linearLeft(x) - laurel = linearRight(laurel) - laurel = postLaurelNorm(laurel) - // Add residual connection - return x + laurel - } -}` +/// Laurel block - uses shared implementation +typealias ${modelName}LaurelBlock = LaurelBlock<${configClass}> +` } /** * Generate AltUp (Alternating Updates) block - * Efficient sparse computation with predict/correct steps + * Uses shared AltUpBlock component. */ -export function generateAltUpBlock(modelName: string, normType: string): string { +export function generateAltUpBlock(modelName: string, configClass: string): string { return `// MARK: - AltUp Block -/// Alternating Updates module for efficient sparse computation -class ${modelName}AltUp: Module { - let numInputs: Int - let activeIdx: Int - let hiddenSize: Int - let altupCoefClip: Float? - - @ModuleInfo(key: "correct_output_scale") var correctOutputScale: MLXArray - @ModuleInfo(key: "correction_coefs") var correctionCoefs: Linear - @ModuleInfo(key: "prediction_coefs") var predictionCoefs: Linear - @ModuleInfo(key: "modality_router") var modalityRouter: Linear - @ModuleInfo(key: "router_norm") var routerNorm: ${normType} - - init(_ config: ${modelName}Configuration) { - self.numInputs = config.altupNumInputs - self.activeIdx = config.altupActiveIdx - self.hiddenSize = config.hiddenSize - self.altupCoefClip = config.altupCoefClip - - _correctOutputScale.wrappedValue = MLXArray.zeros([config.hiddenSize]) - _correctionCoefs.wrappedValue = Linear(numInputs, numInputs, bias: false) - _predictionCoefs.wrappedValue = Linear(numInputs, numInputs * numInputs, bias: false) - _modalityRouter.wrappedValue = Linear(config.hiddenSize, numInputs, bias: false) - _routerNorm.wrappedValue = ${normType}(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func computeRouterModalities(_ x: MLXArray) -> MLXArray { - let routerInputs = routerNorm(x) * pow(Float(hiddenSize), -1.0) - let routed = modalityRouter(routerInputs).asType(.float32) - return tanh(routed) - } - - /// Predict step: modifies input using learned coefficients - /// Input: [numInputs, batch, seq, hidden] -> Output: [numInputs, batch, seq, hidden] - func predict(_ hiddenStates: MLXArray) -> MLXArray { - let modalities = computeRouterModalities(hiddenStates[activeIdx]) - - // Compute prediction coefficients with optional clipping - var weight = predictionCoefs.weight.asType(.float32) - if let clipVal = altupCoefClip { - weight = clip(weight, min: -clipVal, max: clipVal) - } - - // Manual linear: modalities @ weight.T - var allCoefs = matmul(modalities.asType(.float32), weight.T) - let shape = modalities.shape - allCoefs = allCoefs.reshaped([shape[0], shape[1], numInputs, numInputs]) - allCoefs = allCoefs.transposed(0, 1, 3, 2) - - // Convert to float32 for better precision - let xUp = hiddenStates.asType(.float32) - let xPermuted = xUp.transposed(1, 2, 3, 0) - var predictions = matmul(xPermuted, allCoefs) - predictions = predictions.transposed(3, 0, 1, 2) - predictions = predictions + xUp - - return predictions.asType(hiddenStates.dtype) - } - - /// Correct step: refines predictions based on activated output - func correct(_ predictions: MLXArray, activated: MLXArray) -> MLXArray { - let modalities = computeRouterModalities(activated) - - // Compute correction coefficients with optional clipping - var weight = correctionCoefs.weight.asType(.float32) - if let clipVal = altupCoefClip { - weight = clip(weight, min: -clipVal, max: clipVal) - } - - // Manual linear + 1.0: modalities @ weight.T + 1.0 - var allCoefs = matmul(modalities.asType(.float32), weight.T) + 1.0 - let activeX = predictions[activeIdx] - let innovation = activated - activeX - - // allCoefs: [batch, seq, numInputs] -> [numInputs, batch, seq] - allCoefs = allCoefs.transposed(2, 0, 1) - - // innovation: [batch, seq, hidden] - // We need to broadcast: [numInputs, batch, seq, 1] * [1, batch, seq, hidden] - let innovationExpanded = innovation.expandedDimensions(axis: 0) - let allCoefsExpanded = allCoefs.expandedDimensions(axis: -1) - let corrected = innovationExpanded * allCoefsExpanded + predictions - - return corrected.asType(activated.dtype) - } +/// AltUpConfiguration conformance for shared AltUpBlock +extension ${configClass}: AltUpConfiguration {} - func scaleCorrectOutput(_ corrected: MLXArray) -> MLXArray { - return corrected * correctOutputScale - } -}` +/// AltUp block - uses shared implementation +typealias ${modelName}AltUp = AltUpBlock<${configClass}> +` } diff --git a/packages/hf2swift/src/generator/index.ts b/packages/hf2swift/src/generator/index.ts index 1c18d56..ee2265b 100644 --- a/packages/hf2swift/src/generator/index.ts +++ b/packages/hf2swift/src/generator/index.ts @@ -61,43 +61,55 @@ export class SwiftGenerator { private modelName: string private configClass: string private features: ModelFeatures + private configJson?: Record - constructor(modelName: string, features?: ModelFeatures) { + constructor( + modelName: string, + options?: { features?: ModelFeatures; configJson?: Record } + ) { this.modelName = toPascal(modelName) this.configClass = `${this.modelName}Configuration` - this.features = features ?? getModelFeatures(modelName) + this.configJson = options?.configJson + // Features now merge architectural + config values + this.features = options?.features ?? getModelFeatures(modelName, this.configJson) } /** * Generate complete Swift file */ generate(_modules: ParsedModule[], configJson?: Record): string { - const normType = `${this.modelName}RMSNorm` + // Use configJson from generate() or constructor + const effectiveConfig = configJson ?? this.configJson ?? {} + // Re-derive features if configJson provided at generate time + const features = + configJson && !this.configJson + ? getModelFeatures(this.modelName.toLowerCase(), configJson) + : this.features // Always generate config struct (with or without json - defaults are set based on model features) const parts: string[] = [ generateHeader(this.modelName), - generateConfigFromJson(configJson ?? {}, this.modelName, this.features), - generateRmsNorm(this.modelName, this.features), - generateHelpers(this.features) + generateConfigFromJson(effectiveConfig, this.modelName, features), + generateRmsNorm(this.modelName, features), + generateHelpers(features) ] // Add AltUp block if needed (must come before DecoderLayer) - if (this.features.hasAltUp) { - parts.push(generateAltUpBlock(this.modelName, normType)) + if (features.hasAltUp) { + parts.push(generateAltUpBlock(this.modelName, this.configClass)) } // Add Laurel block if needed (must come before DecoderLayer) - if (this.features.hasLaurel) { - parts.push(generateLaurelBlock(this.modelName, this.configClass, normType)) + if (features.hasLaurel) { + parts.push(generateLaurelBlock(this.modelName, this.configClass)) } // Core components - parts.push(generateAttention(this.modelName, this.configClass, this.features)) - parts.push(generateMlp(this.modelName, this.configClass, this.features)) - parts.push(generateDecoderLayer(this.modelName, this.configClass, this.features)) - parts.push(generateModelInner(this.modelName, this.configClass, this.features)) - parts.push(generateModel(this.modelName, this.configClass, this.features)) + parts.push(generateAttention(this.modelName, this.configClass, features)) + parts.push(generateMlp(this.modelName, this.configClass, features)) + parts.push(generateDecoderLayer(this.modelName, this.configClass, features)) + parts.push(generateModelInner(this.modelName, this.configClass, features)) + parts.push(generateModel(this.modelName, this.configClass, features)) const code = parts.filter(Boolean).join("\n\n") return formatSwift(code) diff --git a/packages/hf2swift/src/generator/model-defs/gemma.ts b/packages/hf2swift/src/generator/model-defs/gemma.ts new file mode 100644 index 0000000..665eea3 --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/gemma.ts @@ -0,0 +1,78 @@ +/** + * Gemma model family definition + * + * Includes: Gemma 3, Gemma 3n + * Features: Gemma-style RMSNorm (1+weight), GELU approximate activation, + * Q/K norms, 4 norms per layer, embedding scaling. + * + * Gemma 3n adds: AltUp, Laurel, KV-sharing, sparse activation, VLM support. + */ + +import { DEFAULT_CONFIG, type ModelDefinition, type ArchitecturalFeatures } from "./types.js" + +/** + * Check if model is Gemma 3n (needs special handling) + */ +export function isGemma3n(modelType: string): boolean { + const lower = modelType.toLowerCase() + return lower.includes("gemma3n") || lower.includes("gemma-3n") || lower.includes("gemma_3n") +} + +const gemmaBaseArchitectural: ArchitecturalFeatures = { + rmsNormStyle: "gemma", + activation: "geluApproximate", + useClipResidual: true, + useEmbeddingScale: true, + hasQKNorms: true, + normsPerLayer: 4 +} + +export const gemma3n: ModelDefinition = { + name: "Gemma3n", + + matches: isGemma3n, + + architectural: { + ...gemmaBaseArchitectural, + rmsNormStyle: "standard", // Gemma3n uses standard RMSNorm + useClipResidual: false, + hasAltUp: true, + hasLaurel: true, + hasPerLayerInputs: true, + hasKVSharing: true, + hasPerLayerIntermediateSize: true, + hasSparseActivation: true, + hasVNorm: true, + hasLogitSoftcapping: true, + attentionScale: 1.0 + }, + + configDefaults: { + ...DEFAULT_CONFIG, + useSlidingWindow: true, + ropeTheta: 1000000, + hasLocalRopeTheta: true, + rmsNormEps: 1e-6, + hasWeightTying: true + } +} + +export const gemma3: ModelDefinition = { + name: "Gemma3", + + // Matches "gemma3" but not "gemma3n" + matches: (modelType) => { + const lower = modelType.toLowerCase() + return (lower.includes("gemma3") || lower.includes("gemma-3")) && !isGemma3n(modelType) + }, + + architectural: gemmaBaseArchitectural, + + configDefaults: { + ...DEFAULT_CONFIG, + useSlidingWindow: true, + ropeTheta: 1000000, + hasLocalRopeTheta: true, + rmsNormEps: 1e-6 + } +} diff --git a/packages/hf2swift/src/generator/model-defs/gpt-oss.ts b/packages/hf2swift/src/generator/model-defs/gpt-oss.ts new file mode 100644 index 0000000..fd79166 --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/gpt-oss.ts @@ -0,0 +1,36 @@ +/** + * GPT-OSS model family definition + * + * Mixture of Experts architecture with attention sinks and sliding window. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +export const gptOss: ModelDefinition = { + name: "GPT-OSS", + + matches: (modelType) => { + const lower = modelType.toLowerCase() + return lower.includes("gpt_oss") || lower.includes("gptoss") || lower.includes("gpt-oss") + }, + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu", + hasMoE: true, + hasAttentionSinks: true, + useCustomSwiGLU: true, + useTraditionalRope: true + }, + + configDefaults: { + ...DEFAULT_CONFIG, + useSlidingWindow: true, + ropeTheta: 150000, + hasAttentionBias: true, + hasMlpBias: true, + slidingWindow: 128, + numExperts: 128, + numExpertsPerTok: 4 + } +} diff --git a/packages/hf2swift/src/generator/model-defs/index.ts b/packages/hf2swift/src/generator/model-defs/index.ts new file mode 100644 index 0000000..4a0b15f --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/index.ts @@ -0,0 +1,85 @@ +/** + * Model definitions registry + * + * Central registry of all supported model families. + * Order matters - more specific matchers should come first. + */ + +// Re-export types +export type { + ArchitecturalFeatures, + ConfigValues, + ModelFeatures, + ModelDefinition +} from "./types.js" + +export { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG } from "./types.js" + +// Import model definitions +import { gemma3n, gemma3, isGemma3n } from "./gemma.js" +import { qwen3, qwen2 } from "./qwen.js" +import { mistral3, mistral } from "./mistral.js" +import { phi } from "./phi.js" +import { llama } from "./llama.js" +import { gptOss } from "./gpt-oss.js" +import { smolLm3 } from "./smollm.js" + +import type { ModelDefinition, ArchitecturalFeatures, ConfigValues } from "./types.js" +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG } from "./types.js" + +// Re-export isGemma3n for external use +export { isGemma3n } + +/** + * Model registry - order matters! + * More specific matchers should come first. + */ +const MODEL_REGISTRY: ModelDefinition[] = [ + // Gemma (3n before 3) + gemma3n, + gemma3, + + // Qwen (3 before 2) + qwen3, + qwen2, + + // Mistral (3 before base) + mistral3, + mistral, + + // Others (no ordering needed) + phi, + gptOss, + smolLm3, + llama // Llama is generic, keep last among specific models +] + +/** + * Find matching model definition + */ +export function findModelDefinition(modelType: string): ModelDefinition | undefined { + return MODEL_REGISTRY.find((def) => def.matches(modelType)) +} + +/** + * Get architectural features for a model type + */ +export function getArchitecturalFeatures(modelType: string): ArchitecturalFeatures { + const def = findModelDefinition(modelType) + return def?.architectural ?? DEFAULT_ARCHITECTURAL +} + +/** + * Get default config values for a model type + */ +export function getDefaultConfigValues(modelType: string): ConfigValues { + const def = findModelDefinition(modelType) + return def?.configDefaults ?? DEFAULT_CONFIG +} + +/** + * Get list of all supported model names + */ +export function getSupportedModels(): string[] { + return MODEL_REGISTRY.map((def) => def.name) +} diff --git a/packages/hf2swift/src/generator/model-defs/llama.ts b/packages/hf2swift/src/generator/model-defs/llama.ts new file mode 100644 index 0000000..58a1f05 --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/llama.ts @@ -0,0 +1,23 @@ +/** + * Llama model family definition + * + * Includes: Llama 2, Llama 3, Llama 3.1, Llama 3.2, etc. + * Standard transformer architecture with SiLU activation. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +export const llama: ModelDefinition = { + name: "Llama", + + matches: (modelType) => modelType.toLowerCase().includes("llama"), + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu" + }, + + configDefaults: { + ...DEFAULT_CONFIG + } +} diff --git a/packages/hf2swift/src/generator/model-defs/mistral.ts b/packages/hf2swift/src/generator/model-defs/mistral.ts new file mode 100644 index 0000000..69d3fa2 --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/mistral.ts @@ -0,0 +1,58 @@ +/** + * Mistral model family definition + * + * Includes: Mistral 7B, Mixtral, Mistral 3 (Ministral) + * Mistral 3/Ministral adds YaRN RoPE and removes sliding window. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +/** + * Check if model is Mistral 3 / Ministral 3 + */ +function isMistral3(modelType: string): boolean { + const lower = modelType.toLowerCase() + return ( + lower.includes("mistral3") || + lower.includes("mistral-3") || + lower.includes("ministral3") || + lower.includes("ministral-3") + ) +} + +export const mistral3: ModelDefinition = { + name: "Mistral3", + + matches: isMistral3, + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu", + hasYarnRope: true + }, + + configDefaults: { + ...DEFAULT_CONFIG, + ropeTheta: 1000000 + } +} + +export const mistral: ModelDefinition = { + name: "Mistral", + + // Matches "mistral" or "ministral" but not "mistral3/ministral3" + matches: (modelType) => { + const lower = modelType.toLowerCase() + return (lower.includes("mistral") || lower.includes("ministral")) && !isMistral3(modelType) + }, + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu" + }, + + configDefaults: { + ...DEFAULT_CONFIG, + useSlidingWindow: true + } +} diff --git a/packages/hf2swift/src/generator/model-defs/phi.ts b/packages/hf2swift/src/generator/model-defs/phi.ts new file mode 100644 index 0000000..ed1027d --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/phi.ts @@ -0,0 +1,25 @@ +/** + * Phi model family definition + * + * Includes: Phi-3, Phi-3.5, Phi-4 + * Features: Fused QKV projection, fused gate_up_proj. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +export const phi: ModelDefinition = { + name: "Phi", + + matches: (modelType) => modelType.toLowerCase().includes("phi"), + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu", + hasFusedQKV: true, + hasFusedGateUp: true + }, + + configDefaults: { + ...DEFAULT_CONFIG + } +} diff --git a/packages/hf2swift/src/generator/model-defs/qwen.ts b/packages/hf2swift/src/generator/model-defs/qwen.ts new file mode 100644 index 0000000..b2128e8 --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/qwen.ts @@ -0,0 +1,48 @@ +/** + * Qwen model family definition + * + * Includes: Qwen2, Qwen2.5, Qwen3 + * Qwen3 adds Q/K norms and uses weight tying. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +export const qwen3: ModelDefinition = { + name: "Qwen3", + + matches: (modelType) => modelType.toLowerCase().includes("qwen3"), + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu", + hasQKNorms: true + }, + + configDefaults: { + ...DEFAULT_CONFIG, + ropeTheta: 1000000, + rmsNormEps: 1e-6, + hasWeightTying: true + } +} + +export const qwen2: ModelDefinition = { + name: "Qwen2", + + // Matches "qwen" but not "qwen3" + matches: (modelType) => { + const lower = modelType.toLowerCase() + return lower.includes("qwen") && !lower.includes("qwen3") + }, + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu" + }, + + configDefaults: { + ...DEFAULT_CONFIG, + hasAttentionBias: true, + rmsNormEps: 1e-6 + } +} diff --git a/packages/hf2swift/src/generator/model-defs/smollm.ts b/packages/hf2swift/src/generator/model-defs/smollm.ts new file mode 100644 index 0000000..8f1f29f --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/smollm.ts @@ -0,0 +1,28 @@ +/** + * SmolLM model family definition + * + * SmolLM3 has no-RoPE layers and weight tying. + */ + +import { DEFAULT_ARCHITECTURAL, DEFAULT_CONFIG, type ModelDefinition } from "./types.js" + +export const smolLm3: ModelDefinition = { + name: "SmolLM3", + + matches: (modelType) => { + const lower = modelType.toLowerCase() + return lower.includes("smollm3") || lower.includes("smollm-3") || lower.includes("smollm_3") + }, + + architectural: { + ...DEFAULT_ARCHITECTURAL, + activation: "silu", + hasNoRopeLayers: true + }, + + configDefaults: { + ...DEFAULT_CONFIG, + ropeTheta: 5000000, + hasWeightTying: true + } +} diff --git a/packages/hf2swift/src/generator/model-defs/types.ts b/packages/hf2swift/src/generator/model-defs/types.ts new file mode 100644 index 0000000..b768f6d --- /dev/null +++ b/packages/hf2swift/src/generator/model-defs/types.ts @@ -0,0 +1,144 @@ +/** + * Model definition types + * + * Each model family defines its architectural features and config defaults. + * This enables clean separation and easy addition of new models. + */ + +/** + * Architectural features - determined by model family + * These control which Swift code patterns are generated + */ +export interface ArchitecturalFeatures { + /** RMSNorm style: "gemma" uses (1+weight), "standard" uses weight directly */ + rmsNormStyle: "gemma" | "standard" + + /** Activation function: "gelu", "geluApproximate" (Gemma), or "silu" */ + activation: "gelu" | "geluApproximate" | "silu" + + /** Use clipResidual for float16 overflow protection */ + useClipResidual: boolean + + /** Gemma-style embedding scaling (multiply by sqrt(hiddenSize)) */ + useEmbeddingScale: boolean + + /** Has Q/K norms before attention */ + hasQKNorms: boolean + + /** Number of norms per decoder layer (2 for most, 4 for Gemma3) */ + normsPerLayer: 2 | 4 + + /** Use fused QKV projection instead of separate q_proj, k_proj, v_proj */ + hasFusedQKV?: boolean + + /** Use fused gate_up_proj instead of separate gate_proj, up_proj */ + hasFusedGateUp?: boolean + + /** Uses Mixture of Experts architecture */ + hasMoE?: boolean + + /** Has learnable attention sinks (GPT-OSS) */ + hasAttentionSinks?: boolean + + /** Uses custom SwiGLU activation (alpha=1.702, limit=7.0) */ + useCustomSwiGLU?: boolean + + /** Use traditional RoPE instead of modern */ + useTraditionalRope?: boolean + + // === Advanced Features (Gemma3n) === + hasAltUp?: boolean + hasLaurel?: boolean + hasPerLayerInputs?: boolean + hasKVSharing?: boolean + hasPerLayerIntermediateSize?: boolean + hasSparseActivation?: boolean + hasVNorm?: boolean + hasLogitSoftcapping?: boolean + attentionScale?: number + + // === SmolLM3 / Ministral specific === + hasNoRopeLayers?: boolean + hasYarnRope?: boolean +} + +/** + * Config values - read from config.json with defaults + */ +export interface ConfigValues { + /** Sliding window attention support */ + useSlidingWindow: boolean + + /** RoPE theta */ + ropeTheta: number + + /** Has separate local RoPE theta for sliding window layers */ + hasLocalRopeTheta: boolean + + /** Has attention bias */ + hasAttentionBias: boolean + + /** Has MLP bias */ + hasMlpBias: boolean + + /** RMS norm epsilon */ + rmsNormEps: number + + /** Sliding window size */ + slidingWindow?: number + + /** Number of experts (MoE) */ + numExperts?: number + + /** Experts per token (MoE) */ + numExpertsPerTok?: number + + /** Weight tying (use embed_tokens.weight for lm_head) */ + hasWeightTying?: boolean +} + +/** + * Combined model features = Architectural + Config + */ +export type ModelFeatures = ArchitecturalFeatures & ConfigValues + +/** + * Model definition - architectural features + config defaults + matcher + */ +export interface ModelDefinition { + /** Human-readable name */ + name: string + + /** Check if a model type string matches this definition */ + matches: (modelType: string) => boolean + + /** Architectural features (immutable per model family) */ + architectural: ArchitecturalFeatures + + /** Default config values (fallback when config.json missing) */ + configDefaults: ConfigValues +} + +/** + * Default architectural features (used as base) + */ +export const DEFAULT_ARCHITECTURAL: ArchitecturalFeatures = { + rmsNormStyle: "standard", + activation: "gelu", + useClipResidual: false, + useEmbeddingScale: false, + hasQKNorms: false, + normsPerLayer: 2 +} + +/** + * Default config values (used as base) + */ +export const DEFAULT_CONFIG: ConfigValues = { + useSlidingWindow: false, + ropeTheta: 10000, + hasLocalRopeTheta: false, + hasAttentionBias: false, + hasMlpBias: false, + rmsNormEps: 1e-5 +} diff --git a/packages/hf2swift/src/naming.ts b/packages/hf2swift/src/naming.ts index 6cb8460..6756e33 100644 --- a/packages/hf2swift/src/naming.ts +++ b/packages/hf2swift/src/naming.ts @@ -15,9 +15,13 @@ export function toCamel(name: string): string { */ export function toPascal(name: string): string { // Special cases - if (name.toLowerCase() === "gpt_oss") { + const lower = name.toLowerCase() + if (lower === "gpt_oss" || lower === "gptoss" || lower === "gpt-oss") { return "GptOSS" } + if (lower === "smollm3" || lower === "smol_lm3" || lower === "smol_lm_3") { + return "SmolLM3" + } const parts = name.replace(/-/g, "_").split("_") return parts.map(capitalize).join("") diff --git a/packages/node-mlx/src/cli.ts b/packages/node-mlx/src/cli.ts index 11b9184..ccf4d4c 100644 --- a/packages/node-mlx/src/cli.ts +++ b/packages/node-mlx/src/cli.ts @@ -3,7 +3,7 @@ * * Usage: * mlx # Interactive mode with default model - * mlx --model llama-3.2-1b # Use a specific model + * mlx --model phi4 # Use a specific model * mlx "What is 2+2?" # One-shot query * mlx --list # List available models */ diff --git a/packages/node-mlx/src/index.ts b/packages/node-mlx/src/index.ts index 287a38b..c648da0 100644 --- a/packages/node-mlx/src/index.ts +++ b/packages/node-mlx/src/index.ts @@ -225,22 +225,16 @@ export const RECOMMENDED_MODELS = { "qwen-2.5-1.5b": "Qwen/Qwen2.5-1.5B-Instruct", "qwen-2.5-3b": "Qwen/Qwen2.5-3B-Instruct", - // Phi (Microsoft) - Working with fused QKV and RoPE - phi: "mlx-community/Phi-3.5-mini-instruct-4bit", // Default to 3.5 (smaller, faster to download) + // Phi 4 (Microsoft) - Working with fused QKV and RoPE + phi: "mlx-community/phi-4-4bit", // Phi-4 (14B, highest quality) phi4: "mlx-community/phi-4-4bit", "phi-4": "mlx-community/phi-4-4bit", - "phi-3.5": "mlx-community/Phi-3.5-mini-instruct-4bit", - "phi-3.5-mini": "mlx-community/Phi-3.5-mini-instruct-4bit", - phi3: "mlx-community/Phi-3-mini-4k-instruct-4bit", - "phi-3": "mlx-community/Phi-3-mini-4k-instruct-4bit", - "phi-3-mini": "mlx-community/Phi-3-mini-4k-instruct-4bit", - // Llama 3.2 (Meta) - Requires HuggingFace authentication + // Llama 4 (Meta) - Requires HuggingFace authentication // Note: meta-llama models require accepting license at huggingface.co - llama: "meta-llama/Llama-3.2-1B-Instruct", - "llama-3.2": "meta-llama/Llama-3.2-1B-Instruct", - "llama-3.2-1b": "meta-llama/Llama-3.2-1B-Instruct", - "llama-3.2-3b": "meta-llama/Llama-3.2-3B-Instruct", + llama: "meta-llama/Llama-4-Scout-17B-16E-Instruct", + "llama-4": "meta-llama/Llama-4-Scout-17B-16E-Instruct", + "llama-4-scout": "meta-llama/Llama-4-Scout-17B-16E-Instruct", // Gemma 3 (Google) - Standard transformer architecture with sliding window gemma: "mlx-community/gemma-3-1b-it-4bit", diff --git a/packages/swift/PORTING_DECISIONS.md b/packages/swift/PORTING_DECISIONS.md new file mode 100644 index 0000000..a9633fe --- /dev/null +++ b/packages/swift/PORTING_DECISIONS.md @@ -0,0 +1,234 @@ +# Porting Decisions: mlx-lm (Python) → NodeMLXCore (Swift) + +This document tracks architectural decisions made during the port from Apple's `mlx-lm` Python library to Swift. + +## Source of Truth + +**Decision**: Port directly from `mlx-lm` (Python), not from `mlx-swift-lm` (Swift). + +**Why**: + +- `mlx-lm` is updated more frequently and supports more models +- `mlx-swift-lm` lags behind in features and model support +- Direct Python→Swift porting gives us full control + +**Reference**: https://github.com/ml-explore/mlx-lm/tree/main/mlx_lm/models + +**Current Version**: + +- Git Hash: `7585c142a6be9c9245f4ce61d087839776cb8275` +- Ported: 2026-01-12 + +--- + +## Architecture Overview + +``` +Sources/NodeMLXCore/ +├── generated/ # Auto-generated by hf2swift (DO NOT EDIT) +│ └── models/ # Per-model Swift implementations +├── ported/ # LLM-ported from mlx-lm Python +│ ├── KVCache.swift +│ ├── RoPEUtils.swift +│ ├── SwitchLayers.swift +│ └── GemmaRMSNorm.swift +├── shared/ # Hand-written reusable components +│ ├── Protocols.swift +│ ├── StandardAttention.swift +│ ├── StandardMLP.swift +│ ├── AltUpBlock.swift +│ └── ... +└── (root) # Hand-written integration code + ├── Generate.swift + ├── LLMModel.swift + ├── NodeMLXCore.swift + └── Tokenizer.swift +``` + +### Three-Layer Design + +| Layer | Source | Editing | Purpose | +| -------------- | -------------------- | -------------------- | ---------------------------- | +| **generated/** | `hf2swift` generator | ❌ Never | Model-specific code | +| **ported/** | `mlx-lm` Python | 🔄 Re-port to update | Core MLX infrastructure | +| **shared/** | Hand-written | ✅ Free to edit | Shared components, protocols | +| **root** | Hand-written | ✅ Free to edit | Node.js integration | + +--- + +## Shared Components (shared/) + +Reusable Swift components that reduce generated code by ~70%. + +### Protocols + +| Protocol | Purpose | +| ---------------------------- | ----------------------------------------------------- | +| `BaseModelConfiguration` | Common config properties (hiddenSize, numHeads, etc.) | +| `AttentionConfiguration` | Extends base with attention scale | +| `SlidingWindowConfiguration` | Sliding window support | +| `MoEConfiguration` | Mixture of Experts support | +| `AltUpConfiguration` | Alternating Updates (Gemma3n) | +| `LaurelConfiguration` | Low-rank residual (Gemma3n) | + +### Standard Components + +| Component | Used By | Description | +| ------------------------- | ------------ | -------------------------- | +| `StandardAttention` | Llama, Qwen2 | GQA attention with RoPE | +| `StandardMLP` | Llama, Qwen2 | SwiGLU MLP block | +| `StandardDecoderLayer` | Llama, Qwen2 | Pre-norm decoder | +| `FusedQKVAttention` | Phi3, Phi4 | Fused QKV projection | +| `RMSNorm` | Most models | Standard RMS normalization | + +### Specialized Components + +| Component | Used By | Description | +| ---------------- | ------- | ---------------------------------- | +| `AltUpBlock` | Gemma3n | Alternating Updates sparse compute | +| `LaurelBlock` | Gemma3n | Low-rank residual layer | +| `SparseMLP` | Gemma3n | gelu_topk sparse activation | +| `MoESanitizer` | GPT-OSS | MoE weight transformation | +| `MathUtils` | Various | erfinv, clipResidual, topK | + +--- + +## Ported Components (ported/) + +### KVCache (cache.py → KVCache.swift) + +| Python Class | Swift Class | Notes | +| ------------------------- | ----------------------- | ----------------------- | +| `KVCache` | `StandardKVCache` | Grow-in-place, step=256 | +| `RotatingKVCache` | `RotatingKVCache` | Sliding window + sinks | +| `QuantizedKVCache` | `QuantizedKVCache` | 8-bit quantized | +| `create_causal_mask()` | `createCausalMask()` | Window support | +| `create_attention_mask()` | `createAttentionMask()` | MLXFast mask mode | + +**Not Ported**: BatchKVCache, MambaCache, ChunkedKVCache, CacheList (niche use cases) + +### RoPE Utils (rope_utils.py → RoPEUtils.swift) + +| Python | Swift | Notes | +| ------------------- | ------------------ | ------------------------------ | +| `nn.RoPE` | `StandardRoPE` | RoPEProvider wrapper | +| `Llama3RoPE` | `Llama3RoPE` | Smooth frequency interpolation | +| `YarnRoPE` | `YarnRoPE` | Beta-based correction | +| `SuScaledRoPE` | `SuScaledRoPE` | Long context (longrope) | +| `initialize_rope()` | `initializeRope()` | Factory function | + +**Supported rope_type**: default, linear, llama3, yarn, longrope, mrope + +### SwitchLayers (switch_layers.py → SwitchLayers.swift) + +| Python | Swift | Notes | +| ----------------------- | ----------------------- | -------------------------- | +| `_gather_sort()` | `gatherSort()` | Token sorting for batching | +| `_scatter_unsort()` | `scatterUnsort()` | Restore order | +| `SwitchLinear` | `SwitchLinear` | Expert-specific linear | +| `QuantizedSwitchLinear` | `QuantizedSwitchLinear` | Quantized variant | +| `SwitchGLU` | `SwitchGLU` | Gated linear + experts | +| `swiglu()` | `swiGLU()` | Activation function | + +**GPT-OSS Specific**: `gptOssSwiGLU()` with limit=7.0 clipping + +### GemmaRMSNorm (gemma.py → GemmaRMSNorm.swift) + +| Python | Swift | Notes | +| --------- | -------------- | -------------------- | +| `RMSNorm` | `GemmaRMSNorm` | (1 + weight) scaling | + +--- + +## Generator Architecture (hf2swift) + +The `hf2swift` generator creates Swift model code from HuggingFace patterns. + +### Model Definitions + +``` +packages/hf2swift/src/generator/ +├── model-defs/ # One file per model family +│ ├── types.ts # Interfaces + defaults +│ ├── llama.ts # Llama family +│ ├── qwen.ts # Qwen2, Qwen3 +│ ├── gemma.ts # Gemma3, Gemma3n +│ ├── phi.ts # Phi3, Phi4 +│ ├── mistral.ts # Mistral, Mistral3 +│ ├── gpt-oss.ts # GPT-OSS MoE +│ └── smollm.ts # SmolLM3 +└── components/ # Swift code generators + ├── attention.ts + ├── mlp.ts + ├── decoder-layer.ts + └── model.ts +``` + +### Feature-Based Routing + +The generator uses **feature flags**, not model names, to decide what code to generate: + +```typescript +// Simple models → use shared components +if (canUseSharedStandardAttention(features)) { + return `typealias ${model}Attention = StandardAttention<${config}>` +} + +// Complex models → generate custom code +return generateCustomAttention(model, features) +``` + +### Generated Output + +For simple models (Llama, Qwen2): + +```swift +// MARK: - Attention +typealias LlamaAttention = StandardAttention + +// MARK: - MLP +typealias LlamaMLP = StandardMLP + +// MARK: - Decoder Layer +typealias LlamaDecoderLayer = StandardDecoderLayer +``` + +For complex models (Gemma3n, GPT-OSS): Full custom implementation. + +--- + +## Design Principles + +1. **Focus on popular models**: Llama, Qwen, Phi, Gemma, Mistral, GPT-OSS +2. **Skip niche features**: Batch processing, SSM models, prompt caching +3. **Premium Swift quality**: Protocols, documentation, type safety +4. **Testable components**: Shared code is tested once, used everywhere +5. **Clean separation**: Generated vs. ported vs. hand-written +6. **Feature-driven generation**: Not model-name-driven + +--- + +## Updating Components + +### To update ported code + +1. Check latest mlx-lm commit +2. Download Python source +3. Use `/port-python-to-swift` command in Cursor +4. Update git hash in file header +5. Run tests + +### To add a new model + +1. Create `model-defs/.ts` with features and defaults +2. Register in `model-defs/index.ts` +3. Regenerate: `pnpm hf2swift --model --output ...` +4. Run Swift build and tests + +--- + +## Version History + +| Date | mlx-lm Hash | Changes | +| ---------- | ------------- | ----------------------------------------- | +| 2026-01-12 | `7585c142...` | Initial port: KVCache, RoPE, SwitchLayers | diff --git a/packages/swift/Sources/NodeMLXCore/AttentionUtils.swift b/packages/swift/Sources/NodeMLXCore/AttentionUtils.swift deleted file mode 100644 index 225e416..0000000 --- a/packages/swift/Sources/NodeMLXCore/AttentionUtils.swift +++ /dev/null @@ -1,65 +0,0 @@ -// Copyright © 2024 Apple Inc. -// Adapted for NodeMLXCore - -import Foundation -import MLX -import MLXFast - -/// Attention utilities that match Python mlx-lm's interface -/// -/// This provides a single function that automatically routes to -/// attention based on cache type, matching Python's `scaled_dot_product_attention` - -/// Automatic attention with cache update -/// -/// This function matches Python's `scaled_dot_product_attention` in base.py: -/// - Handles cache updating automatically -/// - Transparent to models - they just call this function -/// -/// **Usage in models:** -/// ```swift -/// let output = attentionWithCacheUpdate( -/// queries: queries, -/// keys: keys, -/// values: values, -/// cache: cache, -/// scale: scale, -/// mask: mask -/// ) -/// ``` -/// -/// - Parameters: -/// - queries: Query tensor [B, nHeads, L, D] -/// - keys: Raw key tensor to be cached [B, nKVHeads, L, D] -/// - values: Raw value tensor to be cached [B, nKVHeads, L, D] -/// - cache: Cache instance (any type) -/// - scale: Attention scale factor -/// - mask: Attention mask -/// - Returns: Attention output [B, nHeads, L, D] -public func attentionWithCacheUpdate( - queries: MLXArray, - keys: MLXArray, - values: MLXArray, - cache: KVCache?, - scale: Float, - mask: MLXFast.ScaledDotProductAttentionMaskMode = .none -) -> MLXArray { - guard let cache else { - return MLXFast.scaledDotProductAttention( - queries: queries, - keys: keys, - values: values, - scale: scale, - mask: mask - ) - } - - let (cachedKeys, cachedValues) = cache.update(keys: keys, values: values) - return MLXFast.scaledDotProductAttention( - queries: queries, - keys: cachedKeys, - values: cachedValues, - scale: scale, - mask: mask - ) -} diff --git a/packages/swift/Sources/NodeMLXCore/Generate.swift b/packages/swift/Sources/NodeMLXCore/Generate.swift index 53777c1..a11c76a 100644 --- a/packages/swift/Sources/NodeMLXCore/Generate.swift +++ b/packages/swift/Sources/NodeMLXCore/Generate.swift @@ -1,146 +1,249 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// Generate.swift -// NodeMLXCore -// -// Token generation with sampling strategies. -// -// Based on patterns from mlx-swift-lm (MIT License, ml-explore). -// See: https://github.com/ml-explore/mlx-swift-lm -// +// Text generation loop for autoregressive language models. import Foundation import MLX -import MLXRandom +import MLXNN -// MARK: - Generation Parameters +// MARK: - Generation Configuration -public struct GenerateParameters: Sendable { - /// Maximum tokens to generate +/// Configuration for text generation. +public struct GenerationConfig { + /// Maximum number of tokens to generate. public var maxTokens: Int - /// Sampling temperature (0 = greedy/argmax) + /// Temperature for sampling (0 = greedy, higher = more random). public var temperature: Float - /// Top-p (nucleus) sampling threshold + /// Top-p nucleus sampling threshold. public var topP: Float - /// Penalty for repeating tokens - public var repetitionPenalty: Float? + /// Repetition penalty (1.0 = no penalty). + public var repetitionPenalty: Float - /// Context size for repetition penalty - public var repetitionContextSize: Int + /// Token IDs that signal end of generation. + public var stopTokens: Set + /// Creates a generation configuration. public init( maxTokens: Int = 256, temperature: Float = 0.7, topP: Float = 0.9, - repetitionPenalty: Float? = nil, - repetitionContextSize: Int = 20 + repetitionPenalty: Float = 1.0, + stopTokens: Set = [] ) { self.maxTokens = maxTokens self.temperature = temperature self.topP = topP self.repetitionPenalty = repetitionPenalty - self.repetitionContextSize = repetitionContextSize + self.stopTokens = stopTokens } } -// MARK: - Sampling Strategies +// MARK: - Token Sampling + +/// Samples the next token from logits. +/// +/// - Parameters: +/// - logits: Model output logits [vocab_size] +/// - temperature: Sampling temperature +/// - topP: Nucleus sampling threshold +/// - Returns: Sampled token ID +public func sampleToken( + logits: MLXArray, + temperature: Float, + topP: Float = 1.0 +) -> Int { + // Greedy decoding for temperature 0 + if temperature == 0 { + return argMax(logits).item(Int.self) + } -/// Sample from logits using argmax (greedy decoding) -public func sampleArgmax(_ logits: MLXArray) -> Int { - let token = argMax(logits, axis: -1) - return token.item(Int.self) -} + // Apply temperature + var scaledLogits = logits / temperature -/// Sample from logits using temperature -public func sampleTemperature(_ logits: MLXArray, temperature: Float) -> Int { - let scaled = logits / MLXArray(temperature) - let probs = softmax(scaled, axis: -1) + // Apply top-p (nucleus) sampling if needed + if topP < 1.0 { + scaledLogits = applyTopP(scaledLogits, topP: topP) + } - // Sample from categorical distribution - let uniform = MLXRandom.uniform(low: 0, high: 1, [1]) - let cumsum = cumsum(probs, axis: -1) - let token = argMax(cumsum .>= uniform, axis: -1) + // Sample from the distribution + let probs = softmax(scaledLogits) + let token = categorical(probs) return token.item(Int.self) } -/// Sample from logits using top-p (nucleus) sampling -public func sampleTopP(_ logits: MLXArray, temperature: Float, topP: Float) -> Int { - // Apply temperature - let scaled = logits / MLXArray(temperature) - let probs = softmax(scaled, axis: -1) +/// Applies top-p (nucleus) sampling by zeroing low-probability tokens. +private func applyTopP(_ logits: MLXArray, topP: Float) -> MLXArray { + let probs = softmax(logits) + let sortedIndices = argSort(probs) + let sortedProbs = probs[sortedIndices] - // Sort probabilities in descending order - let sortedIndices = argSort(probs, axis: -1) - // Reverse to get descending order - let reversedIndices = sortedIndices[.ellipsis, .stride(by: -1)] - let sortedProbs = take(probs, reversedIndices, axis: -1) + // Find cumulative probabilities + let cumProbs = cumsum(sortedProbs) - // Compute cumulative probabilities - let cumProbs = cumsum(sortedProbs, axis: -1) + // Find tokens below threshold + let belowThreshold = cumProbs .<= (1.0 - topP) - // Find cutoff index where cumulative prob exceeds topP - let mask = cumProbs .<= MLXArray(topP) - let numTokens = sum(mask.asType(.int32)).item(Int.self) + 1 + // Mask out tokens below threshold + var result = logits + let maskValue = Float.leastNormalMagnitude + result = which(belowThreshold, MLXArray(maskValue), sortedProbs) - // Keep only top-p tokens - let topIndices = reversedIndices[0 ..< numTokens] - let topProbs = sortedProbs[0 ..< numTokens] + // Unsort back to original order + let unsorted = MLXArray.zeros(like: logits) + unsorted[sortedIndices] = result - // Renormalize - let normalizedProbs = topProbs / sum(topProbs) + return unsorted +} - // Sample from truncated distribution - let uniform = MLXRandom.uniform(low: 0, high: 1, [1]) - let cumsum2 = cumsum(normalizedProbs, axis: -1) - let sampleIdx = argMax(cumsum2 .>= uniform, axis: -1).item(Int.self) +// MARK: - Generation Loop + +/// Generates text from a language model. +/// +/// - Parameters: +/// - model: The language model to use +/// - inputIds: Initial token IDs +/// - config: Generation configuration +/// - onToken: Callback for each generated token +/// - Returns: Array of generated token IDs (excluding input) +public func generate( + model: any LLMModel, + inputIds: [Int], + config: GenerationConfig = GenerationConfig(), + onToken: ((Int) -> Bool)? = nil +) -> [Int] { + var generatedTokens: [Int] = [] + var cache: [KVCacheProtocol]? = model.newCache() + + // Convert input to MLXArray + var currentIds = MLXArray(inputIds.map { Int32($0) }).reshaped([1, inputIds.count]) + + // Process prompt (prefill) + var logits = model(currentIds, cache: &cache) + eval(logits, cache as Any) + + // Get logits for last token + var nextLogits = logits[0..., -1, 0...] + + // Generation loop + for _ in 0 ..< config.maxTokens { + // Sample next token + let nextToken = sampleToken( + logits: nextLogits, + temperature: config.temperature, + topP: config.topP + ) + + // Check for stop token + if config.stopTokens.contains(nextToken) { + break + } + + generatedTokens.append(nextToken) + + // Callback for streaming + if let onToken { + if !onToken(nextToken) { + break + } + } + + // Prepare next input + currentIds = MLXArray([Int32(nextToken)]).reshaped([1, 1]) + + // Generate next logits + logits = model(currentIds, cache: &cache) + eval(logits, cache as Any) + + nextLogits = logits[0..., -1, 0...] + } - return topIndices[sampleIdx].item(Int.self) + return generatedTokens } -/// Main sampling function that dispatches to the right strategy -public func sample(_ logits: MLXArray, params: GenerateParameters) -> Int { - if params.temperature == 0 { - sampleArgmax(logits) - } else if params.topP > 0, params.topP < 1 { - sampleTopP(logits, temperature: params.temperature, topP: params.topP) - } else { - sampleTemperature(logits, temperature: params.temperature) - } +// MARK: - Streaming Generation + +/// Result of a single generation step. +public struct GenerationStep { + /// The generated token ID. + public let tokenId: Int + + /// Whether generation is complete. + public let isComplete: Bool + + /// Decoded text for this token (if decoder provided). + public let text: String? } -// MARK: - Repetition Penalty - -/// Apply repetition penalty to logits -public func applyRepetitionPenalty( - _ logits: MLXArray, - generatedTokens: [Int], - penalty: Float, - contextSize: Int -) -> MLXArray { - guard penalty != 1.0, !generatedTokens.isEmpty else { - return logits +/// Streaming generator for incremental text generation. +public class StreamingGenerator { + private let model: any LLMModel + private let config: GenerationConfig + private var cache: [KVCacheProtocol]? + private var tokenCount: Int = 0 + + /// Creates a streaming generator. + public init(model: any LLMModel, config: GenerationConfig = GenerationConfig()) { + self.model = model + self.config = config } - // Get recent tokens within context window - let recentTokens = Array(generatedTokens.suffix(contextSize)) - guard !recentTokens.isEmpty else { - return logits - } + /// Processes the initial prompt and returns the first token. + public func processPrompt(_ inputIds: [Int]) -> GenerationStep { + cache = model.newCache() - // Create penalty mask - let uniqueTokens = Array(Set(recentTokens)) - let indices = MLXArray(uniqueTokens.map { Int32($0) }) + let currentIds = MLXArray(inputIds.map { Int32($0) }).reshaped([1, inputIds.count]) + let logits = model(currentIds, cache: &cache) + eval(logits, cache as Any) - // Get logits at penalized positions - let selectedLogits = take(logits, indices, axis: -1) + let nextLogits = logits[0..., -1, 0...] + let nextToken = sampleToken( + logits: nextLogits, + temperature: config.temperature, + topP: config.topP + ) - // Apply penalty: divide positive logits, multiply negative - let positiveLogits = maximum(selectedLogits, MLXArray(0)) - let negativeLogits = minimum(selectedLogits, MLXArray(0)) - let penalized = positiveLogits / MLXArray(penalty) + negativeLogits * MLXArray(penalty) + tokenCount = 1 - // Scatter back into original logits using putAlong - return putAlong(logits, indices, values: penalized, axis: -1) + return GenerationStep( + tokenId: nextToken, + isComplete: config.stopTokens.contains(nextToken), + text: nil + ) + } + + /// Generates the next token given the previous one. + public func nextStep(previousToken: Int) -> GenerationStep { + guard tokenCount < config.maxTokens else { + return GenerationStep(tokenId: 0, isComplete: true, text: nil) + } + + let currentIds = MLXArray([Int32(previousToken)]).reshaped([1, 1]) + let logits = model(currentIds, cache: &cache) + eval(logits, cache as Any) + + let nextLogits = logits[0..., -1, 0...] + let nextToken = sampleToken( + logits: nextLogits, + temperature: config.temperature, + topP: config.topP + ) + + tokenCount += 1 + + return GenerationStep( + tokenId: nextToken, + isComplete: config.stopTokens.contains(nextToken) || tokenCount >= config.maxTokens, + text: nil + ) + } + + /// Resets the generator state. + public func reset() { + cache = nil + tokenCount = 0 + } } diff --git a/packages/swift/Sources/NodeMLXCore/KVCache.swift b/packages/swift/Sources/NodeMLXCore/KVCache.swift deleted file mode 100644 index 35db097..0000000 --- a/packages/swift/Sources/NodeMLXCore/KVCache.swift +++ /dev/null @@ -1,383 +0,0 @@ -// Copyright © 2024 Apple Inc. -// Adapted for NodeMLXCore - Core cache functionality from mlx-swift-lm - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - KVCache Protocol - -/// Interface for Key/Value cache for LLMs. -public protocol KVCache: AnyObject { - /// Get the current offset - var offset: Int { get } - - /// Get the current state (keys, values) - used for KV-sharing in Gemma3n - var state: (keys: MLXArray, values: MLXArray)? { get } - - /// Update the cache with new keys and values and return all keys/values - func update(keys: MLXArray, values: MLXArray) -> (MLXArray, MLXArray) - - /// Create an attention mask for this cache - func makeMask( - n: Int, windowSize: Int?, returnArray: Bool - ) -> MLXFast.ScaledDotProductAttentionMaskMode -} - -// MARK: - Causal Mask Creation - -public func createCausalMask( - n: Int, - offset: Int, - windowSize: Int? = nil, - lengths: MLXArray? = nil -) -> MLXArray { - var rinds = MLXArray(Int32(0) ..< Int32(offset + n)) - var linds = offset != 0 ? MLXArray(Int32(offset) ..< Int32(offset + n)) : rinds - linds = linds[0..., .newAxis] - rinds = rinds[.newAxis] - var mask = linds .>= rinds - - if let windowSize { - mask = mask & (linds .< rinds + windowSize) - } - - if var lengths { - lengths = lengths[0..., .newAxis, .newAxis, .newAxis] - mask = mask & (rinds .< lengths) - } - - return mask -} - -// MARK: - Attention Mask Creation - -/// Create an attention mask using the parameters from the KVCache. -public func createAttentionMask(h: MLXArray, cache: KVCache?) -> MLXArray? { - let t = h.dim(1) - if t > 1 { - var offset = 0 - if let c = cache { - offset = c.offset - } - return createCausalMask(n: t, offset: offset) - } - return nil -} - -/// Create an attention mask with explicit window size parameter. -public func createAttentionMask( - h: MLXArray, - cache: KVCache?, - windowSize: Int? = nil, - returnArray: Bool = false -) -> MLXFast.ScaledDotProductAttentionMaskMode { - let n = h.dim(1) - - // Delegate to cache's makeMask if available - if let cache { - return cache.makeMask(n: n, windowSize: windowSize, returnArray: returnArray) - } - - // Fallback for no cache - if n == 1 { - return .none - } - if returnArray || (windowSize != nil && n > windowSize!) { - return .array(createCausalMask(n: n, offset: 0, windowSize: windowSize)) - } - return .causal -} - -// MARK: - KVCacheSimple - -/// Standard KV cache implementation based on Python's KVCache -/// See https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/models/base.py#L11 -public class KVCacheSimple: KVCache { - public private(set) var offset: Int = 0 - var keys: MLXArray? - var values: MLXArray? - public var step = 256 - - /// Get the current state for KV-sharing - public var state: (keys: MLXArray, values: MLXArray)? { - guard let k = keys, let v = values else { return nil } - // Return only the valid portion (up to offset) - return (k[.ellipsis, .. (MLXArray, MLXArray) { - let previous = offset - - let reset = - if let currentKeys = self.keys, (previous + keys.dim(2)) > currentKeys.dim(2) { - true - } else { - self.keys == nil - } - if reset { - let B = keys.dim(0) - let kvHeads = keys.dim(1) - let kHeadDim = keys.dim(3) - let vHeadDim = values.dim(3) - - let nSteps = (step + keys.dim(2) - 1) / step - let kShape = [B, kvHeads, nSteps * step, kHeadDim] - let vShape = [B, kvHeads, nSteps * step, vHeadDim] - let newK = MLXArray.zeros(kShape, dtype: keys.dtype) - let newV = MLXArray.zeros(vShape, dtype: values.dtype) - - if var currentKeys = self.keys, var currentValues = self.values { - if previous % step != 0 { - currentKeys = currentKeys[.ellipsis, .. MLXFast.ScaledDotProductAttentionMaskMode { - // For single token, no mask needed - if n == 1 { - return .none - } - - // For multi-token sequences - if returnArray || (windowSize != nil && n > windowSize!) { - return .array(createCausalMask(n: n, offset: offset, windowSize: windowSize)) - } - - return .causal - } - - public func reset() { - keys = nil - values = nil - offset = 0 - } -} - -// MARK: - RotatingKVCache - -/// Rotating KV cache for sliding window attention -public class RotatingKVCache: KVCache { - public private(set) var offset: Int = 0 - private var keep: Int - private var keys: MLXArray? - private var values: MLXArray? - private var maxCacheSize: Int - private var step: Int - private var idx: Int = 0 - - public var maxSize: Int? { maxCacheSize } - - /// Get the current state for KV-sharing - public var state: (keys: MLXArray, values: MLXArray)? { - guard let k = keys, let v = values else { return nil } - // Return keys/values in temporal order - return (temporalOrder(k), temporalOrder(v)) - } - - public init(maxSize: Int, keep: Int = 0, step: Int = 256) { - maxCacheSize = maxSize - self.keep = keep - self.step = step - } - - private func trim(trimSize: Int, _ array: MLXArray, append: MLXArray? = nil) -> MLXArray { - var toCat: [MLXArray] = [] - if trimSize > 0 { - toCat = [ - array[.ellipsis, .. MLXArray { - // Rearrange the cache into temporal order, slicing off the end if unused - if idx == array.dim(2) { - array - } else if idx < offset { - concatenated( - [ - array[.ellipsis, .. (MLXArray, MLXArray) { - if self.keys == nil { - self.keys = keys - self.values = values - } else { - // Put the keys/values in temporal order to preserve context - self.keys = temporalOrder(self.keys!) - self.values = temporalOrder(self.values!) - idx = self.keys!.dim(2) - - // Allow temporary cache growth during multi-token processing (e.g., prompt prefill). - // The largest size is maxCacheSize + S - 1 to ensure - // every token gets at least maxCacheSize context - let trimSize = idx - maxCacheSize + 1 - self.keys = trim(trimSize: trimSize, self.keys!, append: keys) - self.values = trim(trimSize: trimSize, self.values!, append: values) - } - - offset += keys.dim(2) - idx = self.keys!.dim(2) - - return (self.keys!, self.values!) - } - - private func updateInPlace(keys: MLXArray, values: MLXArray) -> (MLXArray, MLXArray) { - let B = keys.dim(0) - let nKVHeads = keys.dim(1) - let S = keys.dim(2) - let kHeadDim = keys.dim(3) - let vHeadDim = values.dim(3) - let prev = offset - - // May not have hit the max size yet, so potentially keep growing the cache - if self.keys == nil - || (prev >= self.keys!.dim(2) && self.keys!.dim(2) < maxCacheSize) - { - let newSize = min(step, maxCacheSize - prev) - - let kShape = [B, nKVHeads, newSize, kHeadDim] - let vShape = [B, nKVHeads, newSize, vHeadDim] - let newK = MLXArray.zeros(kShape, dtype: keys.dtype) - let newV = MLXArray.zeros(vShape, dtype: values.dtype) - - if let currentKeys = self.keys, let currentValues = self.values { - self.keys = concatenated([currentKeys, newK], axis: 2) - self.values = concatenated([currentValues, newV], axis: 2) - } else { - self.keys = newK - self.values = newV - } - idx = prev - } - - // Trim if needed - let trimSize = self.keys!.dim(2) - maxCacheSize - if trimSize > 0 { - self.keys = trim(trimSize: trimSize, self.keys!) - self.values = trim(trimSize: trimSize, self.values!) - idx = maxCacheSize - } - - // Rotate if we've hit the end - if idx == maxCacheSize { - idx = keep - } - - // Assign - self.keys![.ellipsis, idx ..< (idx + S), 0...] = keys - self.values![.ellipsis, idx ..< (idx + S), 0...] = values - offset += S - idx += S - - // Return the appropriate cache slice - if offset < maxCacheSize { - return ( - self.keys![.ellipsis, .. (MLXArray, MLXArray) { - let result = - if keys.dim(2) == 1 { - updateInPlace(keys: keys, values: values) - } else { - updateConcat(keys: keys, values: values) - } - return result - } - - /// Optimized mask creation for rotating cache with offset capping - public func makeMask( - n: Int, windowSize: Int?, returnArray: Bool - ) -> MLXFast.ScaledDotProductAttentionMaskMode { - if n > 1 { - // Multi-token case - let actualWindowSize = windowSize ?? maxCacheSize - let cappedOffset = min(maxCacheSize - 1, offset) - - // Decide if we need an array mask - if cappedOffset + n > actualWindowSize || returnArray { - return .array( - createCausalMask(n: n, offset: cappedOffset, windowSize: actualWindowSize)) - } - return .causal - } else { - // Single token case (n == 1) - guard let windowSize else { - return .none - } - - // May need a mask when window_size < max_size and cache has wrapped - if offset >= windowSize, maxCacheSize > windowSize { - var currentIdx = idx - if currentIdx >= maxCacheSize { - currentIdx = 0 - } - - let maskSize = offset < maxCacheSize ? offset + 1 : maxCacheSize - let mask = MLXArray(0 ..< Int32(maskSize)) .>= Int32(maskSize - windowSize) - - // Roll the mask to account for rotation - let rolledMask = roll(mask, shift: currentIdx + 1) - - return .array(rolledMask) - } - return .none - } - } -} - -// MARK: - Helper to create cache array for all layers - -/// Create an array of KV caches, one per layer -public func createLayerCaches(numLayers: Int, maxKVSize: Int? = nil) -> [KVCache] { - if let maxKVSize { - (0 ..< numLayers).map { _ in RotatingKVCache(maxSize: maxKVSize, keep: 4) } - } else { - (0 ..< numLayers).map { _ in KVCacheSimple() } - } -} diff --git a/packages/swift/Sources/NodeMLXCore/LLMModel.swift b/packages/swift/Sources/NodeMLXCore/LLMModel.swift index c0784e5..1b9f2ba 100644 --- a/packages/swift/Sources/NodeMLXCore/LLMModel.swift +++ b/packages/swift/Sources/NodeMLXCore/LLMModel.swift @@ -1,229 +1,182 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// LLMModel.swift -// NodeMLXCore -// -// Protocol defining the interface for language models. -// -// Based on patterns from mlx-swift-lm (MIT License, ml-explore). -// See: https://github.com/ml-explore/mlx-swift-lm -// +// Core LLM model protocol and factory for node-mlx. import Foundation import MLX import MLXNN +// MARK: - Type Aliases for Compatibility + +/// Type alias for backward compatibility with generated models. +/// The generated models use KVCache as a protocol/type constraint. +public typealias KVCache = KVCacheProtocol + +/// Simple KV cache - the default implementation used by generated models. +public typealias KVCacheSimple = StandardKVCache + // MARK: - LLM Model Protocol -/// Protocol that all language models must conform to +/// Protocol that all language models must conform to. +/// +/// This defines the common interface for forward passes, caching, +/// and weight loading across all model architectures. public protocol LLMModel: Module { - /// Vocabulary size of the model + /// Vocabulary size for the model var vocabularySize: Int { get } /// Number of transformer layers var numLayers: Int { get } - /// Forward pass with KV cache for efficient generation - func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray + /// Number of key-value heads per layer + var numKVHeads: Int { get } - /// Forward pass without cache (for simple models) - func callAsFunction(_ inputIds: MLXArray) -> MLXArray + /// Dimension of each attention head + var headDim: Int { get } - /// Create a new KV cache for this model - func newCache() -> [KVCache] + /// Forward pass with optional cache + func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCacheProtocol]?) -> MLXArray - /// Whether this model supports KV caching - var supportsCache: Bool { get } + /// Creates a new cache for generation + func newCache() -> [any KVCacheProtocol] - /// Sanitize weight keys during loading (optional override) + /// Sanitizes weight keys for this model architecture func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] } -// MARK: - Default Implementations - -public extension LLMModel { - /// Default sanitize implementation (no-op) - func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - weights - } - - /// Default cache creation - func newCache() -> [KVCache] { - createLayerCaches(numLayers: numLayers) - } - - /// Default: models don't support cache - var supportsCache: Bool { false } - - /// Default cache implementation - falls back to non-cached version - func callAsFunction(_ inputIds: MLXArray, cache _: inout [KVCache]?) -> MLXArray { - // Default: ignore cache and call simple version - callAsFunction(inputIds) - } -} - -// MARK: - Model Registry +// MARK: - Model Architecture Registry /// Supported model architectures public enum ModelArchitecture: String, CaseIterable { case llama + case qwen2 + case qwen3 case phi3 case gemma3 - case gemma3vlm // Gemma 3 with vision case gemma3n - case qwen2 - case qwen3 case mistral - case mistral3 // Ministral 3 / Mistral 3 - case smollm3 // SmolLM 3 - case gptOss // GPT-OSS MoE model - - /// Get architecture from model_type in config.json - public static func from(modelType: String) -> ModelArchitecture? { - let normalized = modelType.lowercased() - .replacingOccurrences(of: "_", with: "") - .replacingOccurrences(of: "-", with: "") - - // Direct matches first (order matters - more specific first) - if normalized == "llama" { return .llama } - if normalized == "phi3" { return .phi3 } - if normalized == "gemma3n" || normalized == "gemma3ntext" { return .gemma3n } - if normalized == "gemma3" || normalized == "gemma3text" { return .gemma3 } - if normalized == "qwen3" { return .qwen3 } // Check qwen3 before qwen2 - if normalized == "qwen2" { return .qwen2 } - if normalized == "mistral3" || normalized == "ministral3" { return .mistral3 } // Check mistral3 before mistral - if normalized == "mistral" { return .mistral } - if normalized == "smollm3" { return .smollm3 } - if normalized == "gptoss" { return .gptOss } - - // Partial matches - for arch in allCases { - let archNormalized = arch.rawValue - .replacingOccurrences(of: "_", with: "") - .replacingOccurrences(of: "-", with: "") - if normalized.contains(archNormalized) { - return arch - } + case mistral3 + case smollm3 + case gptoss = "gpt_oss" + + /// Creates architecture from HuggingFace model_type string. + public init?(modelType: String) { + let normalized = modelType.lowercased().replacingOccurrences(of: "-", with: "_") + + // Try direct match first + if let arch = ModelArchitecture(rawValue: normalized) { + self = arch + return } - return nil - } - /// Check if this is a VLM architecture - public var isVLM: Bool { - switch self { - case .gemma3vlm: - true + // Handle aliases and variations + switch normalized { + case "qwen2.5", "qwen25": + self = .qwen2 + case "llama2", "llama3", "llama3.1", "llama3.2": + self = .llama + case "phi-3", "phi_3": + self = .phi3 + case "gemma-3", "gemma_3": + self = .gemma3 + case "gemma-3n", "gemma_3n": + self = .gemma3n + case "mistral-3", "mistral_3": + self = .mistral3 + case "smollm-3", "smollm_3": + self = .smollm3 + case "gptoss", "gpt-oss": + self = .gptoss default: - false + return nil } } } // MARK: - Model Factory -/// Create a model instance from config and weights +/// Factory for creating model instances from configurations. public enum ModelFactory { - public enum ModelError: Error { - case unsupportedArchitecture(String) - case configLoadFailed(String) - case weightLoadFailed(String) - } - - /// Create model from downloaded directory + /// Creates a model instance from a configuration dictionary. + /// + /// - Parameters: + /// - architecture: The model architecture to create + /// - config: JSON configuration dictionary + /// - Returns: Instantiated model + /// - Throws: DecodingError if configuration is invalid public static func createModel( - modelDirectory: URL, - architecture: ModelArchitecture + architecture: ModelArchitecture, + config: [String: Any] ) throws -> any LLMModel { + let jsonData = try JSONSerialization.data(withJSONObject: config) + let decoder = JSONDecoder() + switch architecture { - case .phi3: - let config = try loadConfig(Phi3Configuration.self, from: modelDirectory) - return Phi3Model(config) case .llama: - let config = try loadConfig(LlamaConfiguration.self, from: modelDirectory) - return LlamaModel(config) - case .gemma3n: - let config = try loadConfig(Gemma3nConfiguration.self, from: modelDirectory) - return Gemma3nModel(config) + let cfg = try decoder.decode(LlamaConfiguration.self, from: jsonData) + return LlamaModel(cfg) + case .qwen2: - let config = try loadConfig(Qwen2Configuration.self, from: modelDirectory) - return Qwen2Model(config) + let cfg = try decoder.decode(Qwen2Configuration.self, from: jsonData) + return Qwen2Model(cfg) + case .qwen3: - let config = try loadConfig(Qwen3Configuration.self, from: modelDirectory) - return Qwen3Model(config) + let cfg = try decoder.decode(Qwen3Configuration.self, from: jsonData) + return Qwen3Model(cfg) + + case .phi3: + let cfg = try decoder.decode(Phi3Configuration.self, from: jsonData) + return Phi3Model(cfg) + case .gemma3: - // Gemma 3 uses standard transformer architecture with some Gemma-specific features - let config = try loadConfig(Gemma3Configuration.self, from: modelDirectory) - return Gemma3Model(config) - case .gemma3vlm: - // Gemma 3 Vision-Language Model - let config = try loadConfig(Gemma3VLMConfiguration.self, from: modelDirectory) - return Gemma3VLMModel(config) - case .mistral: - let config = try loadConfig(MistralConfiguration.self, from: modelDirectory) - return MistralModel(config) - case .mistral3: - let config = try loadConfig(Mistral3Configuration.self, from: modelDirectory) - return Mistral3Model(config) - case .smollm3: - let config = try loadConfig(Smollm3Configuration.self, from: modelDirectory) - return Smollm3Model(config) - case .gptOss: - let config = try loadConfig(GptOSSConfiguration.self, from: modelDirectory) - return GptOSSModel(config) - } - } + let cfg = try decoder.decode(Gemma3Configuration.self, from: jsonData) + return Gemma3Model(cfg) - private static func loadConfig(_: T.Type, from directory: URL) throws -> T { - let configPath = directory.appendingPathComponent("config.json") - let data = try Data(contentsOf: configPath) - return try JSONDecoder().decode(T.self, from: data) - } + case .gemma3n: + let cfg = try decoder.decode(Gemma3nConfiguration.self, from: jsonData) + return Gemma3nModel(cfg) - /// Detect if a model is a VLM by checking for vision_config in config.json - public static func detectVLM(modelDirectory: URL) -> Bool { - let configPath = modelDirectory.appendingPathComponent("config.json") - guard let data = try? Data(contentsOf: configPath), - let json = try? JSONSerialization.jsonObject(with: data) as? [String: Any] - else { - return false - } + case .mistral: + let cfg = try decoder.decode(MistralConfiguration.self, from: jsonData) + return MistralModel(cfg) - // VLM configs have vision_config - return json["vision_config"] != nil - } + case .mistral3: + let cfg = try decoder.decode(Mistral3Configuration.self, from: jsonData) + return Mistral3Model(cfg) - /// Get architecture, automatically detecting VLM - public static func detectArchitecture(modelDirectory: URL) throws -> ModelArchitecture { - let configPath = modelDirectory.appendingPathComponent("config.json") - let data = try Data(contentsOf: configPath) - guard let json = try JSONSerialization.jsonObject(with: data) as? [String: Any], - let modelType = json["model_type"] as? String - else { - throw ModelError.configLoadFailed("Missing model_type in config.json") - } + case .smollm3: + let cfg = try decoder.decode(SmolLM3Configuration.self, from: jsonData) + return SmolLM3Model(cfg) - // Gemma3n needs special handling for config in text_config - if modelType.lowercased().contains("gemma3n") { - return .gemma3n + case .gptoss: + let cfg = try decoder.decode(GptOSSConfiguration.self, from: jsonData) + return GptOSSModel(cfg) } + } - // Check for VLM first - if let visionConfig = json["vision_config"] as? [String: Any] { - // Check if vision is disabled (skip_vision: true indicates text-only quantized from VLM) - let skipVision = visionConfig["skip_vision"] as? Bool ?? false - if !skipVision { - // It's a VLM - check which type - if modelType.lowercased().contains("gemma") { - return .gemma3vlm - } - // Add other VLM types here as needed - } + /// Detects the model architecture from a configuration dictionary. + /// + /// - Parameter config: JSON configuration dictionary + /// - Returns: Detected architecture, or nil if unknown + public static func detectArchitecture(from config: [String: Any]) -> ModelArchitecture? { + // Try model_type field first + if let modelType = config["model_type"] as? String { + return ModelArchitecture(modelType: modelType) } - // Fall back to text-only architecture detection - guard let arch = ModelArchitecture.from(modelType: modelType) else { - throw ModelError.unsupportedArchitecture(modelType) + // Try architectures array + if let architectures = config["architectures"] as? [String], + let first = architectures.first + { + // Parse architecture name (e.g., "LlamaForCausalLM" -> "llama") + let normalized = first + .replacingOccurrences(of: "ForCausalLM", with: "") + .replacingOccurrences(of: "Model", with: "") + .lowercased() + return ModelArchitecture(modelType: normalized) } - return arch + return nil } } diff --git a/packages/swift/Sources/NodeMLXCore/MoELayers.swift b/packages/swift/Sources/NodeMLXCore/MoELayers.swift deleted file mode 100644 index 1436ff7..0000000 --- a/packages/swift/Sources/NodeMLXCore/MoELayers.swift +++ /dev/null @@ -1,300 +0,0 @@ -// -// MoELayers.swift -// NodeMLXCore -// -// Mixture of Experts (MoE) layers for GPT-OSS and similar architectures. -// -// Based on patterns from mlx-lm switch_layers.py and gpt_oss.py: -// - https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/models/switch_layers.py -// - https://github.com/ml-explore/mlx-examples/blob/main/llms/mlx_lm/models/gpt_oss.py -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Custom SwiGLU Activation for GPT-OSS - -/// GPT-OSS uses a modified SwiGLU with specific parameters -/// ```python -/// def swiglu(x_linear, x_glu, alpha=1.702, limit=7.0): -/// x_glu = clip(x_glu, max=limit) -/// x_linear = clip(x_linear, min=-limit, max=limit) -/// glu_scaled = alpha * x_glu -/// sig = sigmoid(glu_scaled) -/// out_glu = x_glu * sig -/// return out_glu * (x_linear + 1) -/// ``` -public func gptOssSwiGLU( - _ xLinear: MLXArray, - _ xGlu: MLXArray, - alpha: Float = 1.702, - limit: Float = 7.0 -) -> MLXArray { - // Clip inputs - let clippedGlu = clip(xGlu, max: MLXArray(limit)) - let clippedLinear = clip(xLinear, min: MLXArray(-limit), max: MLXArray(limit)) - - // Scaled sigmoid gate - let gluScaled = clippedGlu * alpha - let sig = sigmoid(gluScaled) - let outGlu = clippedGlu * sig - - // Apply to linear with +1 bias - return outGlu * (clippedLinear + 1) -} - -// MARK: - Expert Router - -/// Routes tokens to top-k experts based on learned gating -public class MoERouter: Module { - @ModuleInfo(key: "weight") var weight: MLXArray - @ModuleInfo(key: "bias") var bias: MLXArray? - - let hiddenSize: Int - let numExperts: Int - let topK: Int - - public init(hiddenSize: Int, numExperts: Int, topK: Int, bias: Bool = true) { - self.hiddenSize = hiddenSize - self.numExperts = numExperts - self.topK = topK - - _weight.wrappedValue = MLXArray.zeros([numExperts, hiddenSize]) - if bias { - _bias.wrappedValue = MLXArray.zeros([numExperts]) - } else { - _bias.wrappedValue = nil - } - } - - /// Forward pass returns (weights, indices) for top-k experts per token - public func callAsFunction(_ x: MLXArray) -> (weights: MLXArray, indices: MLXArray) { - // x: [batch, seq, hidden] -> logits: [batch, seq, numExperts] - var logits = matmul(x, weight.T) - if let b = bias { - logits = logits + b - } - - // Get top-k experts using argPartition - // argPartition partitions so that the k largest are at the end - let kth = numExperts - topK - let partitionedIndices = argPartition(logits, kth: kth, axis: -1) - - // Take the last topK indices (the largest) - let indices = partitionedIndices[.ellipsis, kth...] - - // Gather the corresponding logit values - let topKLogits = takeAlong(logits, indices, axis: -1) - - // Softmax over selected experts to get weights - let weights = softmax(topKLogits, axis: -1) - - return (weights, indices) - } -} - -// MARK: - SwitchGLU Expert Layer - -/// A single GLU expert with gate, up, and down projections -public class GLUExpert: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - let useCustomSwiGLU: Bool - - public init( - inputDims: Int, - hiddenDims: Int, - bias: Bool = false, - useCustomSwiGLU: Bool = true - ) { - self.useCustomSwiGLU = useCustomSwiGLU - _gateProj.wrappedValue = Linear(inputDims, hiddenDims, bias: bias) - _upProj.wrappedValue = Linear(inputDims, hiddenDims, bias: bias) - _downProj.wrappedValue = Linear(hiddenDims, inputDims, bias: bias) - } - - public func callAsFunction(_ x: MLXArray) -> MLXArray { - let gate = gateProj(x) - let up = upProj(x) - - let hidden: MLXArray = if useCustomSwiGLU { - // GPT-OSS style activation - gptOssSwiGLU(up, gate) - } else { - // Standard SwiGLU - silu(gate) * up - } - - return downProj(hidden) - } -} - -// MARK: - SwitchGLU (Batched Expert MoE) - -/// SwitchGLU implements batched expert computation for MoE layers. -/// -/// This matches the Python mlx-lm SwitchGLU implementation which uses -/// batched operations for efficient expert computation. -/// -/// Weight structure: -/// - experts.gate_proj.weight: [num_experts, hidden_dims, input_dims] -/// - experts.up_proj.weight: [num_experts, hidden_dims, input_dims] -/// - experts.down_proj.weight: [num_experts, input_dims, hidden_dims] -public class SwitchGLU: Module { - @ModuleInfo(key: "gate_proj") var gateProj: MLXArray - @ModuleInfo(key: "up_proj") var upProj: MLXArray - @ModuleInfo(key: "down_proj") var downProj: MLXArray - - // Bias tensors (optional) - @ModuleInfo(key: "gate_proj_bias") var gateProjBias: MLXArray? - @ModuleInfo(key: "up_proj_bias") var upProjBias: MLXArray? - @ModuleInfo(key: "down_proj_bias") var downProjBias: MLXArray? - - let numExperts: Int - let inputDims: Int - let hiddenDims: Int - let useBias: Bool - let useCustomSwiGLU: Bool - - public init( - inputDims: Int, - hiddenDims: Int, - numExperts: Int, - bias: Bool = false, - useCustomSwiGLU: Bool = true - ) { - self.inputDims = inputDims - self.hiddenDims = hiddenDims - self.numExperts = numExperts - useBias = bias - self.useCustomSwiGLU = useCustomSwiGLU - - // Initialize expert weights: [num_experts, out_features, in_features] - _gateProj.wrappedValue = MLXArray.zeros([numExperts, hiddenDims, inputDims]) - _upProj.wrappedValue = MLXArray.zeros([numExperts, hiddenDims, inputDims]) - _downProj.wrappedValue = MLXArray.zeros([numExperts, inputDims, hiddenDims]) - - if bias { - _gateProjBias.wrappedValue = MLXArray.zeros([numExperts, hiddenDims]) - _upProjBias.wrappedValue = MLXArray.zeros([numExperts, hiddenDims]) - _downProjBias.wrappedValue = MLXArray.zeros([numExperts, inputDims]) - } else { - _gateProjBias.wrappedValue = nil - _upProjBias.wrappedValue = nil - _downProjBias.wrappedValue = nil - } - } - - /// Forward pass routes tokens to selected experts and computes weighted output - /// - /// - Parameters: - /// - x: Input tensor [batch * seq, hidden] - /// - indices: Expert indices for each token [batch * seq, topK] - /// - Returns: Expert output [batch * seq, hidden] - public func callAsFunction(_ x: MLXArray, indices: MLXArray) -> MLXArray { - // Get the selected expert weights for each token - // indices: [tokens, topK] -> gather from [numExperts, hiddenDims, inputDims] - let selectedGate = gateProj[indices] // [tokens, topK, hiddenDims, inputDims] - let selectedUp = upProj[indices] - let selectedDown = downProj[indices] // [tokens, topK, inputDims, hiddenDims] - - // Expand x for broadcasting: [tokens, 1, inputDims, 1] - let xExpanded = x[.ellipsis, .newAxis, 0..., .newAxis] - - // Batched matmul: [tokens, topK, hiddenDims, inputDims] @ [tokens, 1, inputDims, 1] - // Result: [tokens, topK, hiddenDims, 1] -> squeeze to [tokens, topK, hiddenDims] - var gateOut = squeezed(matmul(selectedGate, xExpanded), axis: -1) - var upOut = squeezed(matmul(selectedUp, xExpanded), axis: -1) - - // Apply bias if present - if let gateBias = gateProjBias, let upBias = upProjBias { - let selectedGateBias = gateBias[indices] // [tokens, topK, hiddenDims] - let selectedUpBias = upBias[indices] - gateOut = gateOut + selectedGateBias - upOut = upOut + selectedUpBias - } - - // Apply activation - let hidden: MLXArray = if useCustomSwiGLU { - gptOssSwiGLU(upOut, gateOut) - } else { - silu(gateOut) * upOut - } - - // Down projection: [tokens, topK, inputDims, hiddenDims] @ [tokens, topK, hiddenDims, 1] - let hiddenExpanded = hidden[.ellipsis, .newAxis] - var output = squeezed(matmul(selectedDown, hiddenExpanded), axis: -1) // [tokens, topK, inputDims] - - if let downBias = downProjBias { - let selectedDownBias = downBias[indices] - output = output + selectedDownBias - } - - return output - } -} - -// MARK: - Full MoE MLP Layer - -/// Complete MoE MLP layer with router and experts -/// This is used in GPT-OSS style models where each decoder layer has an MoE MLP -public class MoEMLP: Module { - @ModuleInfo(key: "router") var router: MoERouter - @ModuleInfo(key: "experts") var experts: SwitchGLU - - let numExperts: Int - let topK: Int - - public init( - hiddenSize: Int, - intermediateSize: Int, - numExperts: Int, - topK: Int, - bias: Bool = true, - useCustomSwiGLU: Bool = true - ) { - self.numExperts = numExperts - self.topK = topK - - _router.wrappedValue = MoERouter( - hiddenSize: hiddenSize, - numExperts: numExperts, - topK: topK, - bias: bias - ) - _experts.wrappedValue = SwitchGLU( - inputDims: hiddenSize, - hiddenDims: intermediateSize, - numExperts: numExperts, - bias: bias, - useCustomSwiGLU: useCustomSwiGLU - ) - } - - public func callAsFunction(_ x: MLXArray) -> MLXArray { - let shape = x.shape - let batchSeq = shape.dropLast().reduce(1, *) - let hidden = shape.last! - - // Flatten to [batch * seq, hidden] - let xFlat = x.reshaped([batchSeq, hidden]) - - // Get routing weights and expert indices - let (weights, indices) = router(xFlat) - - // Get expert outputs [batch * seq, topK, hidden] - let expertOutput = experts(xFlat, indices: indices) - - // Weighted sum of expert outputs - // weights: [batch * seq, topK] -> [batch * seq, topK, 1] - let weightsExpanded = weights[.ellipsis, .newAxis] - let weightedOutput = sum(expertOutput * weightsExpanded, axis: 1) // [batch * seq, hidden] - - // Reshape back to original shape - return weightedOutput.reshaped(shape) - } -} diff --git a/packages/swift/Sources/NodeMLXCore/ModelLoader.swift b/packages/swift/Sources/NodeMLXCore/ModelLoader.swift deleted file mode 100644 index 15d8473..0000000 --- a/packages/swift/Sources/NodeMLXCore/ModelLoader.swift +++ /dev/null @@ -1,164 +0,0 @@ -// -// ModelLoader.swift -// NodeMLXCore -// -// Downloads and loads MLX models from HuggingFace Hub. -// -// Based on patterns from mlx-swift-lm (MIT License, ml-explore). -// See: https://github.com/ml-explore/mlx-swift-lm -// - -import Foundation -import Hub -import MLX -import MLXNN - -// MARK: - Model Loading Errors - -public enum ModelLoaderError: Error, LocalizedError { - case downloadFailed(String) - case configNotFound(String) - case weightsNotFound(String) - case unsupportedArchitecture(String) - case weightLoadingFailed(String) - - public var errorDescription: String? { - switch self { - case let .downloadFailed(msg): "Download failed: \(msg)" - case let .configNotFound(msg): "Config not found: \(msg)" - case let .weightsNotFound(msg): "Weights not found: \(msg)" - case let .unsupportedArchitecture(msg): "Unsupported architecture: \(msg)" - case let .weightLoadingFailed(msg): "Weight loading failed: \(msg)" - } - } -} - -// MARK: - Model Configuration (from config.json) - -public struct ModelConfig: Codable { - public let modelType: String? - public let hiddenSize: Int? - public let numHiddenLayers: Int? - public let numAttentionHeads: Int? - public let numKeyValueHeads: Int? - public let intermediateSize: Int? - public let vocabSize: Int? - public let maxPositionEmbeddings: Int? - public let ropeTheta: Float? - public let rmsNormEps: Float? - - enum CodingKeys: String, CodingKey { - case modelType = "model_type" - case hiddenSize = "hidden_size" - case numHiddenLayers = "num_hidden_layers" - case numAttentionHeads = "num_attention_heads" - case numKeyValueHeads = "num_key_value_heads" - case intermediateSize = "intermediate_size" - case vocabSize = "vocab_size" - case maxPositionEmbeddings = "max_position_embeddings" - case ropeTheta = "rope_theta" - case rmsNormEps = "rms_norm_eps" - } -} - -// MARK: - Model Loader - -public class ModelLoader { - private let hub: HubApi - - public init() { - hub = HubApi() - } - - /// Download a model from HuggingFace Hub - /// Returns the local directory URL containing the model files - public func download( - modelId: String, - progressHandler: (@Sendable (Progress) -> Void)? = nil - ) async throws -> URL { - let repo = Hub.Repo(id: modelId) - - // Download safetensors and config files - let patterns = ["*.safetensors", "*.json"] - - do { - let modelDir = try await hub.snapshot( - from: repo, - matching: patterns, - progressHandler: progressHandler ?? { _ in } - ) - return modelDir - } catch Hub.HubClientError.authorizationRequired { - throw ModelLoaderError.downloadFailed("Model requires authentication: \(modelId)") - } catch { - throw ModelLoaderError.downloadFailed("\(error)") - } - } - - /// Load configuration from config.json - public func loadConfig(from modelDir: URL) throws -> ModelConfig { - let configURL = modelDir.appendingPathComponent("config.json") - - guard FileManager.default.fileExists(atPath: configURL.path) else { - throw ModelLoaderError.configNotFound(configURL.path) - } - - let data = try Data(contentsOf: configURL) - let config = try JSONDecoder().decode(ModelConfig.self, from: data) - return config - } - - /// Load weights from safetensors files - public func loadWeights(from modelDir: URL) throws -> [String: MLXArray] { - var weights: [String: MLXArray] = [:] - - let enumerator = FileManager.default.enumerator( - at: modelDir, - includingPropertiesForKeys: nil - )! - - for case let url as URL in enumerator { - if url.pathExtension == "safetensors" { - let fileWeights = try loadArrays(url: url) - for (key, value) in fileWeights { - weights[key] = value - } - } - } - - if weights.isEmpty { - throw ModelLoaderError.weightsNotFound(modelDir.path) - } - - return weights - } - - /// Get the model architecture type from config - public func getModelType(from modelDir: URL) throws -> String { - let config = try loadConfig(from: modelDir) - guard let modelType = config.modelType else { - throw ModelLoaderError.configNotFound("model_type not found in config.json") - } - return modelType - } -} - -// MARK: - Weight Utilities - -/// Sanitize weight keys (remove common prefixes, handle quantization) -public func sanitizeWeights(_ weights: [String: MLXArray], prefix: String = "model.") -> [String: MLXArray] { - var sanitized: [String: MLXArray] = [:] - - for (key, value) in weights { - var newKey = key - - // Remove common prefixes - if newKey.hasPrefix(prefix) { - newKey = String(newKey.dropFirst(prefix.count)) - } - - sanitized[newKey] = value - } - - return sanitized -} diff --git a/packages/swift/Sources/NodeMLXCore/Models/LlamaGenerated.swift b/packages/swift/Sources/NodeMLXCore/Models/LlamaGenerated.swift deleted file mode 100644 index 8fd9aae..0000000 --- a/packages/swift/Sources/NodeMLXCore/Models/LlamaGenerated.swift +++ /dev/null @@ -1,326 +0,0 @@ -// -// LlamaGenerated.swift -// NodeMLXCore -// -// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! -// Generated by hf2swift from model patterns. -// Re-run the generator to update this file. -// -// Based on patterns from mlx-lm and mlx-swift-lm. -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Configuration - -public struct LlamaConfiguration: Decodable, Sendable { - public var hiddenSize: Int - public var numHiddenLayers: Int - public var numAttentionHeads: Int - public var numKeyValueHeads: Int - public var intermediateSize: Int - public var vocabSize: Int - public var headDim: Int - public var rmsNormEps: Float - public var ropeTheta: Float - public var maxPositionEmbeddings: Int - public var attentionBias: Bool - public var mlpBias: Bool - public var ropeScaling: [String: StringOrNumber]? - public var modelType: String? - - enum CodingKeys: String, CodingKey { - case textConfig = "text_config" - case hiddenSize = "hidden_size" - case numHiddenLayers = "num_hidden_layers" - case numAttentionHeads = "num_attention_heads" - case numKeyValueHeads = "num_key_value_heads" - case intermediateSize = "intermediate_size" - case vocabSize = "vocab_size" - case headDim = "head_dim" - case rmsNormEps = "rms_norm_eps" - case ropeTheta = "rope_theta" - case maxPositionEmbeddings = "max_position_embeddings" - case attentionBias = "attention_bias" - case mlpBias = "mlp_bias" - case ropeScaling = "rope_scaling" - case modelType = "model_type" - } - - public init(from decoder: Swift.Decoder) throws { - let container = try decoder.container(keyedBy: CodingKeys.self) - - // Helper to decode from text_config or top level - func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { - if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), - let value = try? nested.decode(T.self, forKey: key) - { - return value - } - if let value = try? container.decode(T.self, forKey: key) { - return value - } - if let defaultValue { - return defaultValue - } - throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) - } - - hiddenSize = try decode(.hiddenSize) - numHiddenLayers = try decode(.numHiddenLayers) - numAttentionHeads = try decode(.numAttentionHeads) - numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) - - intermediateSize = try decode(.intermediateSize) - - vocabSize = try decode(.vocabSize) - headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) - ropeTheta = try decode(.ropeTheta, default: 10000.0) - maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) - attentionBias = try decode(.attentionBias, default: false) - mlpBias = try decode(.mlpBias, default: false) - - ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) - modelType = try? container.decode(String.self, forKey: .modelType) - } -} - -// MARK: - RMS Norm - -/// Standard RMSNorm -class LlamaRMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} - -// MARK: - Utility Functions - -// MARK: - Attention - -class LlamaAttention: Module { - @ModuleInfo(key: "q_proj") var qProj: Linear - @ModuleInfo(key: "k_proj") var kProj: Linear - @ModuleInfo(key: "v_proj") var vProj: Linear - @ModuleInfo(key: "o_proj") var oProj: Linear - - let numHeads: Int - let numKVHeads: Int - let headDim: Int - let scale: Float - let rope: RoPE - - init(_ config: LlamaConfiguration) { - numHeads = config.numAttentionHeads - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - scale = 1.0 / sqrt(Float(headDim)) - - let qDim = numHeads * headDim - let kvDim = numKVHeads * headDim - let attnBias = config.attentionBias - - _qProj.wrappedValue = Linear(config.hiddenSize, qDim, bias: attnBias) - _kProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _vProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: attnBias) - rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - - var queries = qProj(hiddenStates).reshaped([B, L, numHeads, headDim]) - var keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - var values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - - // Transpose for attention: [B, heads, L, headDim] - queries = queries.transposed(0, 2, 1, 3) - keys = keys.transposed(0, 2, 1, 3) - values = values.transposed(0, 2, 1, 3) - - // Apply RoPE with cache offset - let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) - - // Update cache - if let c = cache { - (keys, values) = c.update(keys: keys, values: values) - } - - // Attention using MLXFast (handles GQA automatically) - let output = MLXFast.scaledDotProductAttention( - queries: queries, - keys: keys, - values: values, - scale: scale, - mask: mask - ) - - // Reshape back: [B, heads, L, headDim] -> [B, L, hidden] - let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) - return oProj(outputReshaped) - } -} - -// MARK: - MLP - -class LlamaMLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: LlamaConfiguration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} - -// MARK: - Decoder Layer - -class LlamaDecoderLayer: Module { - @ModuleInfo(key: "self_attn") var selfAttn: LlamaAttention - @ModuleInfo(key: "mlp") var mlp: LlamaMLP - @ModuleInfo(key: "input_layernorm") var inputLayernorm: LlamaRMSNorm - @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: LlamaRMSNorm - - init(_ config: LlamaConfiguration, layerIdx _: Int = 0) { - _selfAttn.wrappedValue = LlamaAttention(config) - _mlp.wrappedValue = LlamaMLP(config) - _inputLayernorm.wrappedValue = LlamaRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - _postAttentionLayernorm.wrappedValue = LlamaRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - // 1. Pre-norm + Self-attention - let normed = inputLayernorm(hiddenStates) - let attnOut = selfAttn(normed, mask: mask, cache: &cache) - var h = hiddenStates + attnOut - - // 2. Pre-norm + MLP - let mlpNormed = postAttentionLayernorm(h) - let mlpOut = mlp(mlpNormed) - h = h + mlpOut - return h - } -} - -// MARK: - Model Inner - -class LlamaModelInner: Module { - @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding - @ModuleInfo(key: "layers") var layers: [LlamaDecoderLayer] - @ModuleInfo(key: "norm") var norm: LlamaRMSNorm - - let numLayers: Int - let hiddenSize: Int - - init(_ config: LlamaConfiguration) { - numLayers = config.numHiddenLayers - hiddenSize = config.hiddenSize - - _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) - _layers.wrappedValue = (0 ..< numLayers).map { idx in LlamaDecoderLayer(config, layerIdx: idx) } - _norm.wrappedValue = LlamaRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { - var hiddenStates = embedTokens(inputIds) - let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil) - for i in 0 ..< layers.count { - hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) - } - return norm(hiddenStates) - } -} - -// MARK: - Top-Level Model - -public class LlamaModel: Module, LLMModel { - public let vocabularySize: Int - public let numLayers: Int - public let numKVHeads: Int - public let headDim: Int - - @ModuleInfo(key: "model") var model: LlamaModelInner - @ModuleInfo(key: "lm_head") var lmHead: Linear - - private let config: LlamaConfiguration - - public var supportsCache: Bool { true } - - public init(_ config: LlamaConfiguration) { - self.config = config - vocabularySize = config.vocabSize - numLayers = config.numHiddenLayers - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - _model.wrappedValue = LlamaModelInner(config) - _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) - } - - public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { - var cache: [KVCache?] = Array(repeating: nil, count: numLayers) - let h = model(inputIds, cache: &cache) - return lmHead(h) - } - - public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { - var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } - else { Array(repeating: nil, count: numLayers) } - let h = model(inputIds, cache: &layerCaches) - cache = layerCaches.compactMap(\.self) - return lmHead(h) - } - - public func newCache() -> [KVCache] { - (0 ..< numLayers).map { _ in KVCacheSimple() } - } - - public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result - } -} diff --git a/packages/swift/Sources/NodeMLXCore/Models/Mistral3Generated.swift b/packages/swift/Sources/NodeMLXCore/Models/Mistral3Generated.swift deleted file mode 100644 index c8cfec9..0000000 --- a/packages/swift/Sources/NodeMLXCore/Models/Mistral3Generated.swift +++ /dev/null @@ -1,358 +0,0 @@ -// -// Mistral3Generated.swift -// NodeMLXCore -// -// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! -// Generated by hf2swift from model patterns. -// Re-run the generator to update this file. -// -// Based on patterns from mlx-lm and mlx-swift-lm. -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Configuration - -public struct Mistral3Configuration: Decodable, Sendable { - public var hiddenSize: Int - public var numHiddenLayers: Int - public var numAttentionHeads: Int - public var numKeyValueHeads: Int - public var intermediateSize: Int - public var vocabSize: Int - public var headDim: Int - public var rmsNormEps: Float - public var ropeTheta: Float - public var maxPositionEmbeddings: Int - public var attentionBias: Bool - public var mlpBias: Bool - public var slidingWindow: Int - public var slidingWindowPattern: Int - public var ropeScaling: [String: StringOrNumber]? - public var modelType: String? - - /// Check if a layer is a global attention layer - public func isGlobalLayer(_ layerIdx: Int) -> Bool { - (layerIdx % slidingWindowPattern) == (slidingWindowPattern - 1) - } - - enum CodingKeys: String, CodingKey { - case textConfig = "text_config" - case hiddenSize = "hidden_size" - case numHiddenLayers = "num_hidden_layers" - case numAttentionHeads = "num_attention_heads" - case numKeyValueHeads = "num_key_value_heads" - case intermediateSize = "intermediate_size" - case vocabSize = "vocab_size" - case headDim = "head_dim" - case rmsNormEps = "rms_norm_eps" - case ropeTheta = "rope_theta" - case maxPositionEmbeddings = "max_position_embeddings" - case attentionBias = "attention_bias" - case mlpBias = "mlp_bias" - case slidingWindow = "sliding_window" - case slidingWindowPattern = "sliding_window_pattern" - case ropeScaling = "rope_scaling" - case modelType = "model_type" - } - - public init(from decoder: Swift.Decoder) throws { - let container = try decoder.container(keyedBy: CodingKeys.self) - - // Helper to decode from text_config or top level - func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { - if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), - let value = try? nested.decode(T.self, forKey: key) - { - return value - } - if let value = try? container.decode(T.self, forKey: key) { - return value - } - if let defaultValue { - return defaultValue - } - throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) - } - - hiddenSize = try decode(.hiddenSize) - numHiddenLayers = try decode(.numHiddenLayers) - numAttentionHeads = try decode(.numAttentionHeads) - numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) - - intermediateSize = try decode(.intermediateSize) - - vocabSize = try decode(.vocabSize) - headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) - ropeTheta = try decode(.ropeTheta, default: 10000.0) - maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) - attentionBias = try decode(.attentionBias, default: false) - mlpBias = try decode(.mlpBias, default: false) - - slidingWindow = try decode(.slidingWindow, default: 512) - slidingWindowPattern = try decode(.slidingWindowPattern, default: 6) - ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) - modelType = try? container.decode(String.self, forKey: .modelType) - } -} - -// MARK: - RMS Norm - -/// Standard RMSNorm -class Mistral3RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} - -// MARK: - Utility Functions - -// MARK: - Attention - -class Mistral3Attention: Module { - @ModuleInfo(key: "q_proj") var qProj: Linear - @ModuleInfo(key: "k_proj") var kProj: Linear - @ModuleInfo(key: "v_proj") var vProj: Linear - @ModuleInfo(key: "o_proj") var oProj: Linear - - let numHeads: Int - let numKVHeads: Int - let headDim: Int - let scale: Float - let rope: RoPE - let isSliding: Bool - - init(_ config: Mistral3Configuration, layerIdx: Int) { - numHeads = config.numAttentionHeads - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - scale = 1.0 / sqrt(Float(headDim)) - - let qDim = numHeads * headDim - let kvDim = numKVHeads * headDim - let attnBias = config.attentionBias - - _qProj.wrappedValue = Linear(config.hiddenSize, qDim, bias: attnBias) - _kProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _vProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: attnBias) - isSliding = !config.isGlobalLayer(layerIdx) - let ropeBase = config.ropeTheta - rope = RoPE(dimensions: headDim, traditional: false, base: ropeBase) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - - var queries = qProj(hiddenStates).reshaped([B, L, numHeads, headDim]) - var keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - var values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - - // Transpose for attention: [B, heads, L, headDim] - queries = queries.transposed(0, 2, 1, 3) - keys = keys.transposed(0, 2, 1, 3) - values = values.transposed(0, 2, 1, 3) - - // Apply RoPE with cache offset - let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) - - // Update cache - if let c = cache { - (keys, values) = c.update(keys: keys, values: values) - } - - // Attention using MLXFast (handles GQA automatically) - let output = MLXFast.scaledDotProductAttention( - queries: queries, - keys: keys, - values: values, - scale: scale, - mask: mask - ) - - // Reshape back: [B, heads, L, headDim] -> [B, L, hidden] - let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) - return oProj(outputReshaped) - } -} - -// MARK: - MLP - -class Mistral3MLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: Mistral3Configuration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} - -// MARK: - Decoder Layer - -class Mistral3DecoderLayer: Module { - @ModuleInfo(key: "self_attn") var selfAttn: Mistral3Attention - @ModuleInfo(key: "mlp") var mlp: Mistral3MLP - @ModuleInfo(key: "input_layernorm") var inputLayernorm: Mistral3RMSNorm - @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: Mistral3RMSNorm - - init(_ config: Mistral3Configuration, layerIdx: Int) { - _selfAttn.wrappedValue = Mistral3Attention(config, layerIdx: layerIdx) - _mlp.wrappedValue = Mistral3MLP(config) - _inputLayernorm.wrappedValue = Mistral3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - _postAttentionLayernorm.wrappedValue = Mistral3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - // 1. Pre-norm + Self-attention - let normed = inputLayernorm(hiddenStates) - let attnOut = selfAttn(normed, mask: mask, cache: &cache) - var h = hiddenStates + attnOut - - // 2. Pre-norm + MLP - let mlpNormed = postAttentionLayernorm(h) - let mlpOut = mlp(mlpNormed) - h = h + mlpOut - return h - } -} - -// MARK: - Model Inner - -class Mistral3ModelInner: Module { - @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding - @ModuleInfo(key: "layers") var layers: [Mistral3DecoderLayer] - @ModuleInfo(key: "norm") var norm: Mistral3RMSNorm - - let numLayers: Int - let hiddenSize: Int - let slidingWindow: Int - let slidingWindowPattern: Int - - init(_ config: Mistral3Configuration) { - numLayers = config.numHiddenLayers - hiddenSize = config.hiddenSize - slidingWindow = config.slidingWindow - slidingWindowPattern = config.slidingWindowPattern - _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) - _layers.wrappedValue = (0 ..< numLayers).map { idx in Mistral3DecoderLayer(config, layerIdx: idx) } - _norm.wrappedValue = Mistral3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { - var hiddenStates = embedTokens(inputIds) - let globalLayerIdx = slidingWindowPattern - 1 - let globalCache = globalLayerIdx < cache.count ? cache[globalLayerIdx] : nil - let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) - let slidingMask: MLXFast.ScaledDotProductAttentionMaskMode - if slidingWindowPattern > 1 { - let firstCache = cache.first ?? nil - slidingMask = createAttentionMask(h: hiddenStates, cache: firstCache, windowSize: slidingWindow) - } else { - slidingMask = globalMask - } - for i in 0 ..< layers.count { - let isGlobal = (i % slidingWindowPattern) == (slidingWindowPattern - 1) - let mask = isGlobal ? globalMask : slidingMask - hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) - } - return norm(hiddenStates) - } -} - -// MARK: - Top-Level Model - -public class Mistral3Model: Module, LLMModel { - public let vocabularySize: Int - public let numLayers: Int - public let numKVHeads: Int - public let headDim: Int - - @ModuleInfo(key: "model") var model: Mistral3ModelInner - @ModuleInfo(key: "lm_head") var lmHead: Linear - - private let config: Mistral3Configuration - - public var supportsCache: Bool { true } - - public init(_ config: Mistral3Configuration) { - self.config = config - vocabularySize = config.vocabSize - numLayers = config.numHiddenLayers - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - _model.wrappedValue = Mistral3ModelInner(config) - _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) - } - - public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { - var cache: [KVCache?] = Array(repeating: nil, count: numLayers) - let h = model(inputIds, cache: &cache) - return lmHead(h) - } - - public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { - var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } - else { Array(repeating: nil, count: numLayers) } - let h = model(inputIds, cache: &layerCaches) - cache = layerCaches.compactMap(\.self) - return lmHead(h) - } - - public func newCache() -> [KVCache] { - (0 ..< numLayers).map { i in - let isGlobal = (i % config.slidingWindowPattern) == (config.slidingWindowPattern - 1) - if isGlobal { return KVCacheSimple() } - else { return RotatingKVCache(maxSize: config.slidingWindow, keep: 0) } - } - } - - public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result - } -} diff --git a/packages/swift/Sources/NodeMLXCore/Models/Qwen2Generated.swift b/packages/swift/Sources/NodeMLXCore/Models/Qwen2Generated.swift deleted file mode 100644 index 31eb3b9..0000000 --- a/packages/swift/Sources/NodeMLXCore/Models/Qwen2Generated.swift +++ /dev/null @@ -1,326 +0,0 @@ -// -// Qwen2Generated.swift -// NodeMLXCore -// -// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! -// Generated by hf2swift from model patterns. -// Re-run the generator to update this file. -// -// Based on patterns from mlx-lm and mlx-swift-lm. -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Configuration - -public struct Qwen2Configuration: Decodable, Sendable { - public var hiddenSize: Int - public var numHiddenLayers: Int - public var numAttentionHeads: Int - public var numKeyValueHeads: Int - public var intermediateSize: Int - public var vocabSize: Int - public var headDim: Int - public var rmsNormEps: Float - public var ropeTheta: Float - public var maxPositionEmbeddings: Int - public var attentionBias: Bool - public var mlpBias: Bool - public var ropeScaling: [String: StringOrNumber]? - public var modelType: String? - - enum CodingKeys: String, CodingKey { - case textConfig = "text_config" - case hiddenSize = "hidden_size" - case numHiddenLayers = "num_hidden_layers" - case numAttentionHeads = "num_attention_heads" - case numKeyValueHeads = "num_key_value_heads" - case intermediateSize = "intermediate_size" - case vocabSize = "vocab_size" - case headDim = "head_dim" - case rmsNormEps = "rms_norm_eps" - case ropeTheta = "rope_theta" - case maxPositionEmbeddings = "max_position_embeddings" - case attentionBias = "attention_bias" - case mlpBias = "mlp_bias" - case ropeScaling = "rope_scaling" - case modelType = "model_type" - } - - public init(from decoder: Swift.Decoder) throws { - let container = try decoder.container(keyedBy: CodingKeys.self) - - // Helper to decode from text_config or top level - func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { - if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), - let value = try? nested.decode(T.self, forKey: key) - { - return value - } - if let value = try? container.decode(T.self, forKey: key) { - return value - } - if let defaultValue { - return defaultValue - } - throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) - } - - hiddenSize = try decode(.hiddenSize) - numHiddenLayers = try decode(.numHiddenLayers) - numAttentionHeads = try decode(.numAttentionHeads) - numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) - - intermediateSize = try decode(.intermediateSize) - - vocabSize = try decode(.vocabSize) - headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) - ropeTheta = try decode(.ropeTheta, default: 10000.0) - maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) - attentionBias = try decode(.attentionBias, default: true) - mlpBias = try decode(.mlpBias, default: false) - - ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) - modelType = try? container.decode(String.self, forKey: .modelType) - } -} - -// MARK: - RMS Norm - -/// Standard RMSNorm -class Qwen2RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} - -// MARK: - Utility Functions - -// MARK: - Attention - -class Qwen2Attention: Module { - @ModuleInfo(key: "q_proj") var qProj: Linear - @ModuleInfo(key: "k_proj") var kProj: Linear - @ModuleInfo(key: "v_proj") var vProj: Linear - @ModuleInfo(key: "o_proj") var oProj: Linear - - let numHeads: Int - let numKVHeads: Int - let headDim: Int - let scale: Float - let rope: RoPE - - init(_ config: Qwen2Configuration) { - numHeads = config.numAttentionHeads - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - scale = 1.0 / sqrt(Float(headDim)) - - let qDim = numHeads * headDim - let kvDim = numKVHeads * headDim - let attnBias = config.attentionBias - - _qProj.wrappedValue = Linear(config.hiddenSize, qDim, bias: attnBias) - _kProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _vProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) - _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: attnBias) - rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - - var queries = qProj(hiddenStates).reshaped([B, L, numHeads, headDim]) - var keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - var values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) - - // Transpose for attention: [B, heads, L, headDim] - queries = queries.transposed(0, 2, 1, 3) - keys = keys.transposed(0, 2, 1, 3) - values = values.transposed(0, 2, 1, 3) - - // Apply RoPE with cache offset - let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) - - // Update cache - if let c = cache { - (keys, values) = c.update(keys: keys, values: values) - } - - // Attention using MLXFast (handles GQA automatically) - let output = MLXFast.scaledDotProductAttention( - queries: queries, - keys: keys, - values: values, - scale: scale, - mask: mask - ) - - // Reshape back: [B, heads, L, headDim] -> [B, L, hidden] - let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) - return oProj(outputReshaped) - } -} - -// MARK: - MLP - -class Qwen2MLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: Qwen2Configuration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} - -// MARK: - Decoder Layer - -class Qwen2DecoderLayer: Module { - @ModuleInfo(key: "self_attn") var selfAttn: Qwen2Attention - @ModuleInfo(key: "mlp") var mlp: Qwen2MLP - @ModuleInfo(key: "input_layernorm") var inputLayernorm: Qwen2RMSNorm - @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: Qwen2RMSNorm - - init(_ config: Qwen2Configuration, layerIdx _: Int = 0) { - _selfAttn.wrappedValue = Qwen2Attention(config) - _mlp.wrappedValue = Qwen2MLP(config) - _inputLayernorm.wrappedValue = Qwen2RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - _postAttentionLayernorm.wrappedValue = Qwen2RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - // 1. Pre-norm + Self-attention - let normed = inputLayernorm(hiddenStates) - let attnOut = selfAttn(normed, mask: mask, cache: &cache) - var h = hiddenStates + attnOut - - // 2. Pre-norm + MLP - let mlpNormed = postAttentionLayernorm(h) - let mlpOut = mlp(mlpNormed) - h = h + mlpOut - return h - } -} - -// MARK: - Model Inner - -class Qwen2ModelInner: Module { - @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding - @ModuleInfo(key: "layers") var layers: [Qwen2DecoderLayer] - @ModuleInfo(key: "norm") var norm: Qwen2RMSNorm - - let numLayers: Int - let hiddenSize: Int - - init(_ config: Qwen2Configuration) { - numLayers = config.numHiddenLayers - hiddenSize = config.hiddenSize - - _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) - _layers.wrappedValue = (0 ..< numLayers).map { idx in Qwen2DecoderLayer(config, layerIdx: idx) } - _norm.wrappedValue = Qwen2RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { - var hiddenStates = embedTokens(inputIds) - let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil) - for i in 0 ..< layers.count { - hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) - } - return norm(hiddenStates) - } -} - -// MARK: - Top-Level Model - -public class Qwen2Model: Module, LLMModel { - public let vocabularySize: Int - public let numLayers: Int - public let numKVHeads: Int - public let headDim: Int - - @ModuleInfo(key: "model") var model: Qwen2ModelInner - @ModuleInfo(key: "lm_head") var lmHead: Linear - - private let config: Qwen2Configuration - - public var supportsCache: Bool { true } - - public init(_ config: Qwen2Configuration) { - self.config = config - vocabularySize = config.vocabSize - numLayers = config.numHiddenLayers - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - _model.wrappedValue = Qwen2ModelInner(config) - _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) - } - - public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { - var cache: [KVCache?] = Array(repeating: nil, count: numLayers) - let h = model(inputIds, cache: &cache) - return lmHead(h) - } - - public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { - var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } - else { Array(repeating: nil, count: numLayers) } - let h = model(inputIds, cache: &layerCaches) - cache = layerCaches.compactMap(\.self) - return lmHead(h) - } - - public func newCache() -> [KVCache] { - (0 ..< numLayers).map { _ in KVCacheSimple() } - } - - public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result - } -} diff --git a/packages/swift/Sources/NodeMLXCore/NodeMLXCore.swift b/packages/swift/Sources/NodeMLXCore/NodeMLXCore.swift index 4f99f4e..835b7a5 100644 --- a/packages/swift/Sources/NodeMLXCore/NodeMLXCore.swift +++ b/packages/swift/Sources/NodeMLXCore/NodeMLXCore.swift @@ -1,595 +1,339 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// NodeMLXCore.swift -// NodeMLXCore -// -// Main entry point for LLM inference without mlx-swift-lm dependency. -// -// Copyright © 2026 Sebastian Software GmbH. All rights reserved. -// +// Core Swift implementation for node-mlx. +// Provides the main integration point between Node.js and MLX. import Foundation import Hub import MLX -import MLXFast import MLXNN -import MLXRandom -import Tokenizers -// MARK: - Public API +// MARK: - Generation Result + +/// Result of text generation. +public struct GenerationResult: Sendable { + /// The generated text. + public let text: String + + /// Number of tokens generated. + public let tokenCount: Int + + /// Tokens per second. + public let tokensPerSecond: Float + + /// Time to first token in seconds. + public let timeToFirstToken: Double + + /// Total generation time in seconds. + public let totalTime: Double +} -/// Main interface for LLM operations +// MARK: - LLM Engine + +/// Main engine for loading and running language models. +/// +/// This class manages model loading, tokenization, and generation, +/// providing a high-level API for the Node.js bindings. public class LLMEngine { private var model: (any LLMModel)? - private var vlmModel: Gemma3VLMModel? // VLM-specific reference private var tokenizer: HFTokenizer? - private var modelDirectory: URL? - private var imageProcessor: ImageProcessor? - private var _isVLM: Bool = false - private var _isGemma: Bool = false // For enforcing Gemma chat template + private var modelPath: String? + + /// Whether a model is currently loaded. + public var isLoaded: Bool { model != nil } - /// Whether the loaded model is a Vision-Language Model - public var isVLM: Bool { _isVLM } + /// Whether this is a vision-language model (VLM). + public var isVLM: Bool { false } // Not implemented yet + /// Creates an empty engine. public init() {} - // MARK: - Model Loading - - /// Load a model from HuggingFace Hub - public func loadModel( - modelId: String, - progressHandler: ((Float) -> Void)? = nil - ) async throws { - // Download model files - let hub = HubApi() - let repo = Hub.Repo(id: modelId) - - let directory = try await hub.snapshot( - from: repo, - matching: ["*.safetensors", "*.json", "tokenizer*", "vocab*", "merges*"], - progressHandler: { progress in - progressHandler?(Float(progress.fractionCompleted)) - } - ) + /// Loads a model from HuggingFace Hub or local directory. + /// + /// - Parameter modelId: HuggingFace model ID or local path + /// - Throws: Error if model cannot be loaded + public func loadModel(modelId: String) async throws { + // Check if it's a local path + let fileManager = FileManager.default + if fileManager.fileExists(atPath: modelId) { + try await loadModelFromPath(modelId) + } else { + // Download from HuggingFace Hub + let hubApi = HubApi() + let repo = Hub.Repo(id: modelId) + let localPath = try await hubApi.snapshot(from: repo, matching: ["*.json", "*.safetensors"]) + try await loadModelFromPath(localPath.path) + } + } - modelDirectory = directory + /// Loads a model from a local directory. + /// + /// - Parameter path: Path to model directory containing config.json and weights + /// - Throws: Error if model cannot be loaded + private func loadModelFromPath(_ path: String) async throws { + let url = URL(fileURLWithPath: path) - // Detect architecture from config - let configPath = directory.appendingPathComponent("config.json") + // Load configuration + let configPath = url.appendingPathComponent("config.json") let configData = try Data(contentsOf: configPath) - let configDict = try JSONSerialization.jsonObject(with: configData) as? [String: Any] ?? [:] - - guard configDict["model_type"] as? String != nil else { - throw LLMEngineError.invalidConfig("model_type not found in config.json") + guard let config = try JSONSerialization.jsonObject(with: configData) as? [String: Any] else { + throw LLMEngineError.invalidConfig("Cannot parse config.json") } - // Detect architecture (including VLM detection) - let architecture = try ModelFactory.detectArchitecture(modelDirectory: directory) - - // Track if this is a VLM - _isVLM = architecture.isVLM - - // Track if this is a Gemma model (for enforcing chat template) - _isGemma = architecture == .gemma3 || architecture == .gemma3vlm || architecture == .gemma3n + // Detect architecture + guard let architecture = ModelFactory.detectArchitecture(from: config) else { + let modelType = config["model_type"] as? String ?? "unknown" + throw LLMEngineError.unsupportedModel("Unsupported model type: \(modelType)") + } // Create model - let model = try ModelFactory.createModel( - modelDirectory: directory, - architecture: architecture - ) + let newModel = try ModelFactory.createModel(architecture: architecture, config: config) - // Keep VLM-specific reference for image generation - if let vlm = model as? Gemma3VLMModel { - vlmModel = vlm - // Create image processor for VLM - imageProcessor = ImageProcessor(config: .siglip) - } + // Load weights + let weights = try loadWeights(from: url, config: config) - // Load weights first - let weights = try loadWeights(from: directory) - let sanitizedWeights = model.sanitize(weights: weights) + // Sanitize weight keys + let sanitizedWeights = newModel.sanitize(weights: weights) - // Quantize if needed - use dynamic quantization based on weight presence - if let quantizationConfig = configDict["quantization"] as? [String: Any], - let groupSize = quantizationConfig["group_size"] as? Int, - let bits = quantizationConfig["bits"] as? Int + // Handle quantization + if let quantConfig = config["quantization"] as? [String: Any], + let groupSize = quantConfig["group_size"] as? Int, + let bits = quantConfig["bits"] as? Int { - // Quantize modules that have .scales weights - // The filter returns (groupSize, bits, mode) if the module should be quantized - quantize(model: model) { path, _ in - if sanitizedWeights["\(path).scales"] != nil { - (groupSize, bits, .affine) - } else { - nil + quantize(model: newModel, predicate: { weightPath, _ in + // Check if this weight has quantization scales + if sanitizedWeights["\(weightPath).scales"] != nil { + return (groupSize, bits, .affine) } - } + return nil + }) } - // Apply weights to model - model.update(parameters: ModuleParameters.unflattened(sanitizedWeights)) - - // Force evaluation of weights to ensure they're loaded to GPU - eval(model) - - self.model = model + // Apply weights + newModel.update(parameters: ModuleParameters.unflattened(sanitizedWeights)) + eval(newModel.parameters()) // Load tokenizer - tokenizer = try await HFTokenizer(modelDirectory: directory) - } + let newTokenizer = try await HFTokenizer(path: path) - // MARK: - Generation + model = newModel + tokenizer = newTokenizer + modelPath = path + } - /// Generate text from a prompt + /// Generates text from a prompt. + /// + /// - Parameters: + /// - prompt: Input text + /// - config: Generation configuration + /// - onToken: Optional callback for streaming tokens + /// - Returns: Generated text public func generate( prompt: String, - maxTokens: Int = 256, - temperature: Float = 0.7, - topP: Float = 0.9, - repetitionPenalty: Float? = nil, - repetitionContextSize: Int = 20 - ) throws -> GenerationResult { - guard let model else { + config: GenerationConfig = GenerationConfig(), + onToken: ((String) -> Bool)? = nil + ) throws -> String { + guard let model, let tokenizer else { throw LLMEngineError.modelNotLoaded } - guard let tokenizer else { - throw LLMEngineError.tokenizerNotLoaded - } - // Apply chat template to format prompt correctly for the model - var inputTokens: [Int] - if _isGemma { - // Gemma models (including Gemma3n) need explicit chat template formatting - // Some variants don't have chat_template in tokenizer_config.json - let formattedPrompt = "user\n\(prompt)\nmodel\n" - inputTokens = tokenizer.encode(formattedPrompt) - } else { - do { - inputTokens = try tokenizer.applyChatTemplate(userMessage: prompt) - } catch { - // Fallback: use raw prompt - inputTokens = tokenizer.encode(prompt) - } - } - var inputArray = MLXArray(inputTokens.map { Int32($0) }) - inputArray = inputArray.expandedDimensions(axis: 0) // Add batch dimension - - let startTime = Date() - var generatedTokens: [Int] = [] - - // Create KV cache for efficient generation - var cache: [KVCache]? = model.newCache() + // Encode prompt + let inputIds = tokenizer.encode(text: prompt) - // Process prompt (prefill) - all tokens at once - var logits = model(inputArray, cache: &cache) - var lastLogits = logits[0, logits.dim(1) - 1] - eval(lastLogits) - - // Apply repetition penalty if configured - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: inputTokens, penalty: penalty, contextSize: repetitionContextSize) + // Set up stop tokens + var genConfig = config + if let eosId = tokenizer.eosTokenId { + genConfig.stopTokens.insert(eosId) } - // Sample first token - var nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - - // Generation loop - one token at a time with cached context - for _ in 0 ..< maxTokens { - // Check for EOS (both and for chat models) - if let eosId = tokenizer.eosTokenId, nextToken == eosId { - break - } - // Gemma models use (106) for chat - if nextToken == 106 { - break - } - - generatedTokens.append(nextToken) - - // Prepare next input - just the single new token - inputArray = MLXArray([Int32(nextToken)]).expandedDimensions(axis: 0) - - // Forward pass with cache - only processes new token - logits = model(inputArray, cache: &cache) - lastLogits = logits[0, 0] // Single token output - - // Async eval for pipelining - eval(lastLogits) - - // Apply repetition penalty before sampling - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: generatedTokens, penalty: penalty, contextSize: repetitionContextSize) + // Generate tokens + let generatedIds = NodeMLXCore.generate( + model: model, + inputIds: inputIds, + config: genConfig, + onToken: onToken.map { callback in + { tokenId in + let text = tokenizer.decode(tokens: [tokenId]) + return callback(text) + } } - - // Sample next token - nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - } - - let endTime = Date() - let duration = Float(endTime.timeIntervalSince(startTime)) - let tokensPerSecond = duration > 0 ? Float(generatedTokens.count) / duration : 0 - - // Decode generated tokens, skipping special tokens like <|end|> - let generatedText = tokenizer.decode(generatedTokens, skipSpecialTokens: true) - - return GenerationResult( - text: generatedText, - tokenCount: generatedTokens.count, - tokensPerSecond: tokensPerSecond ) + + // Decode result + return tokenizer.decode(tokens: generatedIds) } - /// Generate text with streaming callback + /// Generates text with streaming and returns detailed result. + /// + /// - Parameters: + /// - prompt: Input text + /// - maxTokens: Maximum tokens to generate + /// - temperature: Sampling temperature + /// - topP: Nucleus sampling threshold + /// - repetitionPenalty: Penalty for repeated tokens (optional) + /// - repetitionContextSize: Context size for repetition penalty + /// - onToken: Callback for each generated token + /// - Returns: Generation result with timing information public func generateStream( prompt: String, - maxTokens: Int = 256, - temperature: Float = 0.7, - topP: Float = 0.9, + maxTokens: Int, + temperature: Float, + topP: Float, repetitionPenalty: Float? = nil, - repetitionContextSize: Int = 20, - onToken: @escaping (String) -> Bool // Return false to stop + repetitionContextSize _: Int = 20, + onToken: @escaping (String) -> Bool ) throws -> GenerationResult { - guard let model else { + guard let model, let tokenizer else { throw LLMEngineError.modelNotLoaded } - guard let tokenizer else { - throw LLMEngineError.tokenizerNotLoaded - } - - // Apply chat template to format prompt correctly for the model - var inputTokens: [Int] - if _isGemma { - // Gemma models (including Gemma3n) need explicit chat template formatting - let formattedPrompt = "user\n\(prompt)\nmodel\n" - inputTokens = tokenizer.encode(formattedPrompt) - } else { - do { - inputTokens = try tokenizer.applyChatTemplate(userMessage: prompt) - } catch { - // Fallback: use raw prompt - inputTokens = tokenizer.encode(prompt) - } - } - var inputArray = MLXArray(inputTokens.map { Int32($0) }) - inputArray = inputArray.expandedDimensions(axis: 0) - let startTime = Date() - var generatedTokens: [Int] = [] + let startTime = CFAbsoluteTimeGetCurrent() + var firstTokenTime: CFAbsoluteTime? - // Create KV cache for efficient generation - var cache: [KVCache]? = model.newCache() + // Encode prompt + let inputIds = tokenizer.encode(text: prompt) - // Process prompt (prefill) - all tokens at once - var logits = model(inputArray, cache: &cache) - var lastLogits = logits[0, logits.dim(1) - 1] - eval(lastLogits) - - // Apply repetition penalty if configured - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: inputTokens, penalty: penalty, contextSize: repetitionContextSize) + // Set up config + var config = GenerationConfig( + maxTokens: maxTokens, + temperature: temperature, + topP: topP, + repetitionPenalty: repetitionPenalty ?? 1.0 + ) + if let eosId = tokenizer.eosTokenId { + config.stopTokens.insert(eosId) } - // Sample first token - var nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - - // Generation loop with KV cache - for _ in 0 ..< maxTokens { - // Check for EOS (both and for chat models) - if let eosId = tokenizer.eosTokenId, nextToken == eosId { - break - } - // Gemma models use (106) for chat - if nextToken == 106 { - break - } - - generatedTokens.append(nextToken) - - // Stream the token (skip special tokens like <|end|>) - let tokenText = tokenizer.decode([nextToken], skipSpecialTokens: true) - if !tokenText.isEmpty, !onToken(tokenText) { - break // User requested stop - } - - // Prepare next input - just the single new token - inputArray = MLXArray([Int32(nextToken)]).expandedDimensions(axis: 0) - - // Forward pass with cache - only processes new token - logits = model(inputArray, cache: &cache) - lastLogits = logits[0, 0] // Single token output - eval(lastLogits) - - // Apply repetition penalty before sampling - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: generatedTokens, penalty: penalty, contextSize: repetitionContextSize) + // Generate tokens + let generatedIds = NodeMLXCore.generate( + model: model, + inputIds: inputIds, + config: config, + onToken: { tokenId in + if firstTokenTime == nil { + firstTokenTime = CFAbsoluteTimeGetCurrent() + } + let text = tokenizer.decode(tokens: [tokenId]) + return onToken(text) } + ) - nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - } - - let endTime = Date() - let duration = Float(endTime.timeIntervalSince(startTime)) - let tokensPerSecond = duration > 0 ? Float(generatedTokens.count) / duration : 0 + let endTime = CFAbsoluteTimeGetCurrent() + let totalTime = endTime - startTime + let timeToFirst = (firstTokenTime ?? endTime) - startTime return GenerationResult( - text: tokenizer.decode(generatedTokens, skipSpecialTokens: true), - tokenCount: generatedTokens.count, - tokensPerSecond: tokensPerSecond + text: tokenizer.decode(tokens: generatedIds), + tokenCount: generatedIds.count, + tokensPerSecond: generatedIds.count > 0 ? Float(generatedIds.count) / Float(totalTime) : 0, + timeToFirstToken: timeToFirst, + totalTime: totalTime ) } - // MARK: - VLM Generation - - /// Generate text with image input (for VLMs) + /// Generates text with an image (VLM). + /// + /// - Note: VLM support is not yet implemented. public func generateStreamWithImage( - prompt: String, - imagePath: String, - maxTokens: Int = 256, - temperature: Float = 0.7, - topP: Float = 0.9, - repetitionPenalty: Float? = nil, - repetitionContextSize: Int = 20, - onToken: @escaping (String) -> Bool + prompt _: String, + imagePath _: String, + maxTokens _: Int, + temperature _: Float, + topP _: Float, + repetitionPenalty _: Float? = nil, + repetitionContextSize _: Int = 20, + onToken _: @escaping (String) -> Bool ) throws -> GenerationResult { - guard let vlmModel else { - throw LLMEngineError.notAVLM - } - guard let tokenizer else { - throw LLMEngineError.tokenizerNotLoaded - } - guard let imageProcessor else { - throw LLMEngineError.imageProcessingFailed("No image processor available") - } - - // Load and preprocess image - let pixelValues: MLXArray - do { - pixelValues = try imageProcessor.loadAndPreprocess(path: imagePath) - } catch { - throw LLMEngineError.imageProcessingFailed("Failed to load image: \(error.localizedDescription)") - } - - // For VLM, we need to include the image token ID directly - // The tokenizer doesn't recognize as a special token, so we insert it manually - // Gemma 3 VLM image token ID is 262144 - let imageTokenId = 262_144 - - // First tokenize the prompt without image - var inputTokens: [Int] - do { - inputTokens = try tokenizer.applyChatTemplate(userMessage: prompt) - } catch { - // Fallback: manually construct a VLM-style prompt - let manualPrompt = "user\n\(prompt)\nmodel\n" - inputTokens = tokenizer.encode(manualPrompt) - } - - // Find position after "user\n" to insert image token - // The format is: user\n[IMAGE_HERE]promptmodel\n - // Token IDs: 2 (bos), 105 (start_of_turn), user tokens, 107 (newline) - var insertPos = 0 - for (i, token) in inputTokens.enumerated() { - // Look for the newline token (107) after user - if token == 107, i > 2 { - insertPos = i + 1 - break - } - } - - // Insert image token at the found position - if insertPos > 0, insertPos < inputTokens.count { - inputTokens.insert(imageTokenId, at: insertPos) - } else { - // Fallback: insert after BOS token - inputTokens.insert(imageTokenId, at: 1) - } - - var inputArray = MLXArray(inputTokens.map { Int32($0) }) - inputArray = inputArray.expandedDimensions(axis: 0) - - let startTime = Date() - var generatedTokens: [Int] = [] - - // Create KV cache - var cache: [KVCache]? = vlmModel.newCache() - - // Process prompt with image (prefill) - var logits = vlmModel(inputArray, pixelValues: pixelValues, cache: &cache) - var lastLogits = logits[0, logits.dim(1) - 1] - eval(lastLogits) - - // Apply repetition penalty if configured - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: inputTokens, penalty: penalty, contextSize: repetitionContextSize) - } - - // Sample first token - var nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - - // Generation loop - for _ in 0 ..< maxTokens { - if let eosId = tokenizer.eosTokenId, nextToken == eosId { - break - } - if nextToken == 106 { // - break - } - - generatedTokens.append(nextToken) - - let tokenText = tokenizer.decode([nextToken], skipSpecialTokens: true) - if !tokenText.isEmpty, !onToken(tokenText) { - break - } - - inputArray = MLXArray([Int32(nextToken)]).expandedDimensions(axis: 0) - - // Forward without image (already encoded in KV cache) - logits = vlmModel(inputArray, pixelValues: nil, cache: &cache) - lastLogits = logits[0, 0] - eval(lastLogits) - - // Apply repetition penalty before sampling - if let penalty = repetitionPenalty { - lastLogits = applyRepetitionPenalty(lastLogits, generatedTokens: generatedTokens, penalty: penalty, contextSize: repetitionContextSize) - } - - nextToken = sampleToken(logits: lastLogits, temperature: temperature, topP: topP) - } - - let endTime = Date() - let duration = Float(endTime.timeIntervalSince(startTime)) - let tokensPerSecond = duration > 0 ? Float(generatedTokens.count) / duration : 0 - - return GenerationResult( - text: tokenizer.decode(generatedTokens, skipSpecialTokens: true), - tokenCount: generatedTokens.count, - tokensPerSecond: tokensPerSecond - ) + throw LLMEngineError.unsupportedModel("VLM support not yet implemented") } - // MARK: - Cleanup - - /// Unload the model from memory + /// Unloads the current model. public func unload() { model = nil - vlmModel = nil tokenizer = nil - modelDirectory = nil - imageProcessor = nil - _isVLM = false + modelPath = nil } +} - // MARK: - Private Helpers +// MARK: - Weight Loading - private func loadWeights(from directory: URL) throws -> [String: MLXArray] { - var weights: [String: MLXArray] = [:] +/// Loads model weights from a directory. +/// +/// Supports both safetensors and npz formats. +private func loadWeights(from url: URL, config _: [String: Any]) throws -> [String: MLXArray] { + // Find weight files + let fileManager = FileManager.default + let contents = try fileManager.contentsOfDirectory(at: url, includingPropertiesForKeys: nil) - let enumerator = FileManager.default.enumerator( - at: directory, - includingPropertiesForKeys: nil - )! + // Prefer safetensors + let safetensorFiles = contents.filter { $0.pathExtension == "safetensors" } + let npzFiles = contents.filter { $0.pathExtension == "npz" } - for case let url as URL in enumerator { - if url.pathExtension == "safetensors" { - let fileWeights = try loadArrays(url: url) - for (key, value) in fileWeights { - weights[key] = value - } + var weights: [String: MLXArray] = [:] + + if !safetensorFiles.isEmpty { + // Load all safetensor files + for file in safetensorFiles.sorted(by: { $0.lastPathComponent < $1.lastPathComponent }) { + let fileWeights = try MLX.loadArrays(url: file) + for (key, value) in fileWeights { + weights[key] = value } } - - if weights.isEmpty { - throw LLMEngineError.weightsNotFound + } else if !npzFiles.isEmpty { + // Load first npz file + if let npzFile = npzFiles.first { + weights = try MLX.loadArrays(url: npzFile) } - - return weights + } else { + throw LLMEngineError.weightsNotFound } - private func sampleToken(logits: MLXArray, temperature: Float, topP: Float) -> Int { - if temperature == 0 { - // Greedy decoding - no randomness - let token = argMax(logits, axis: -1) - eval(token) - return Int(token.item(Int32.self)) - } - - // Temperature scaling and convert to probabilities - let temp = MLXArray(temperature) - var logitsFloat = logits - if logitsFloat.dtype == .bfloat16 { - logitsFloat = logitsFloat.asType(.float32) - } - let probs = softmax(logitsFloat / temp, axis: -1) - - // For top-p sampling, use the mlx-swift-lm approach - if topP > 0, topP < 1 { - let topPArray = MLXArray(topP) - - // Sort in ascending order (lowest first) - let sortedIndices = argSort(probs, axis: -1) - let sortedProbs = take(probs, sortedIndices, axis: -1) - - // Cumulative sum (from lowest to highest) - let cumulativeProbs = cumsum(sortedProbs, axis: -1) - - // Keep only tokens where cumulative prob > (1 - topP) - // This keeps the top-p highest probability tokens - let topProbs = MLX.where( - cumulativeProbs .> (1 - topPArray), - sortedProbs, - MLXArray.zeros(like: sortedProbs) - ) - - // Sample using log probabilities (avoid numerical issues) - let sortedToken = MLXRandom.categorical(log(topProbs)) - eval(sortedToken) - - // Map back to original index - let originalIdx = sortedIndices[Int(sortedToken.item(Int32.self))] - eval(originalIdx) - return Int(originalIdx.item(Int32.self)) - } - - // Simple temperature sampling without top-p - let token = MLXRandom.categorical(probs) - eval(token) - return Int(token.item(Int32.self)) - } + return weights } -// MARK: - Types - -public struct GenerationResult { - public let text: String - public let tokenCount: Int - public let tokensPerSecond: Float - - public init(text: String, tokenCount: Int, tokensPerSecond: Float) { - self.text = text - self.tokenCount = tokenCount - self.tokensPerSecond = tokensPerSecond - } -} +// MARK: - Error Types +/// Errors that can occur during LLM engine operations. public enum LLMEngineError: Error, LocalizedError { case modelNotLoaded - case tokenizerNotLoaded case invalidConfig(String) case unsupportedModel(String) case weightsNotFound - case notAVLM - case imageProcessingFailed(String) + case generationFailed(String) public var errorDescription: String? { switch self { case .modelNotLoaded: - "No model loaded. Call loadModel() first." - case .tokenizerNotLoaded: - "No tokenizer loaded." + "No model is loaded" case let .invalidConfig(msg): - "Invalid config: \(msg)" + "Invalid configuration: \(msg)" case let .unsupportedModel(msg): "Unsupported model: \(msg)" case .weightsNotFound: - "No weights found in model directory." - case .notAVLM: - "Model does not support images (not a VLM). Use a vision model like google/gemma-3-4b-it." - case let .imageProcessingFailed(msg): - "Image processing failed: \(msg)" + "No weight files found in model directory" + case let .generationFailed(msg): + "Generation failed: \(msg)" } } } -// MARK: - Convenience - -/// Quick generation without managing engine lifecycle -public func quickGenerate( - modelId: String, - prompt: String, - maxTokens: Int = 256 -) async throws -> String { - let engine = LLMEngine() - try await engine.loadModel(modelId: modelId) - let result = try engine.generate(prompt: prompt, maxTokens: maxTokens) - engine.unload() - return result.text +// MARK: - Quantization Helper + +/// Quantizes model layers that have corresponding scale weights. +private func quantize( + model: Module, + predicate: (String, Module) -> (Int, Int, QuantizationMode)? +) { + model.update(modules: ModuleChildren.unflattened( + model.leafModules().flattened().compactMap { path, module in + guard let (groupSize, bits, mode) = predicate(path, module) else { + return nil + } + if let linear = module as? Linear { + return (path, QuantizedLinear(linear, groupSize: groupSize, bits: bits, mode: mode)) + } + return nil + } + )) } diff --git a/packages/swift/Sources/NodeMLXCore/README.md b/packages/swift/Sources/NodeMLXCore/README.md new file mode 100644 index 0000000..d97cd22 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/README.md @@ -0,0 +1,69 @@ +# NodeMLXCore + +Swift implementation of MLX-based language model inference for Node.js. + +## Architecture + +``` +NodeMLXCore/ +├── generated/ # Auto-generated model code (DO NOT EDIT) +│ └── models/ # One Swift file per model +├── ported/ # Code ported from mlx-lm Python +│ ├── KVCache.swift # KV cache implementations +│ ├── RoPEUtils.swift # Rotary position embeddings +│ └── ... +├── shared/ # Reusable Swift components +│ ├── Protocols.swift # Base configuration protocols +│ ├── Standard*.swift # Generic model components +│ └── ... +└── (root) # Hand-written integration code + ├── Generate.swift # Text generation + ├── LLMModel.swift # Model protocol + ├── NodeMLXCore.swift # C-interface bridge + └── Tokenizer.swift # Tokenization +``` + +## Three-Layer Design + +| Directory | Source | Edit Policy | Purpose | +| ------------ | --------------- | --------------- | ------------------------------ | +| `generated/` | `hf2swift` | ❌ Never edit | Model-specific implementations | +| `ported/` | `mlx-lm` Python | 🔄 Re-port only | Core MLX infrastructure | +| `shared/` | Hand-written | ✅ Free to edit | Reusable components | +| Root files | Hand-written | ✅ Free to edit | Node.js integration | + +## Supported Models + +| Model | Type | Features | +| ------------ | -------------- | -------------------------------- | +| Llama 3.x | Standard | Uses shared components | +| Qwen2, Qwen3 | Standard | Qwen3 has Q/K norms | +| Phi-3, Phi-4 | Fused QKV | Fused projections | +| Gemma3 | 4-norm | Gemma-style RMSNorm | +| Gemma3n | VLM | AltUp, Laurel, sparse activation | +| Mistral | Sliding window | Window attention | +| GPT-OSS | MoE | Mixture of Experts | +| SmolLM3 | No-RoPE layers | Selective RoPE | + +## Quick Start + +### Regenerate a Model + +```bash +pnpm hf2swift --model llama --output packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift +``` + +### Build and Test + +```bash +cd packages/swift +swift build -c release +swift test +``` + +## Documentation + +- **[PORTING_DECISIONS.md](../../PORTING_DECISIONS.md)** - Architectural decisions +- **[generated/README.md](generated/README.md)** - Generated code guidelines +- **[ported/README.md](ported/README.md)** - Porting process +- **[shared/README.md](shared/README.md)** - Shared component catalog diff --git a/packages/swift/Sources/NodeMLXCore/RoPEUtils.swift b/packages/swift/Sources/NodeMLXCore/RoPEUtils.swift deleted file mode 100644 index 7867c19..0000000 --- a/packages/swift/Sources/NodeMLXCore/RoPEUtils.swift +++ /dev/null @@ -1,380 +0,0 @@ -// Copyright © 2024 Apple Inc. -// Adapted for NodeMLXCore - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - RoPE Protocol - -/// Protocol for all RoPE variants to enable polymorphic usage -public protocol RoPEProvider { - func apply(_ x: MLXArray, offset: Int) -> MLXArray -} - -extension RoPE: RoPEProvider { - public func apply(_ x: MLXArray, offset: Int) -> MLXArray { - callAsFunction(x, offset: offset) - } -} - -// MARK: - Llama3RoPE - -public class Llama3RoPE: Module, RoPEProvider { - let dims: Int - let maxPositionEmbeddings: Int - let traditional: Bool - let freqs: MLXArray - - public init( - dims: Int, - maxPositionEmbeddings: Int = 2048, - traditional: Bool = false, - base: Float = 10000, - scalingConfig: [String: StringOrNumber]? = nil - ) { - self.dims = dims - self.maxPositionEmbeddings = maxPositionEmbeddings - self.traditional = traditional - - guard let scalingConfig else { - fatalError("Llama3RoPE requires scaling_config") - } - - let factor = scalingConfig["factor"]?.asFloat() ?? 1.0 - let lowFreqFactor = scalingConfig["low_freq_factor"]?.asFloat() ?? 1.0 - let highFreqFactor = scalingConfig["high_freq_factor"]?.asFloat() ?? 4.0 - let oldContextLen = scalingConfig["original_max_position_embeddings"]?.asFloat() ?? 8192.0 - - let lowFreqWavelen = oldContextLen / lowFreqFactor - let highFreqWavelen = oldContextLen / highFreqFactor - - let indices = MLXArray(stride(from: 0, to: dims, by: 2)) - var frequencies = MLX.pow(base, indices / Float(dims)) - let wavelens = 2 * Float.pi * frequencies - - frequencies = MLX.where( - wavelens .> MLXArray(lowFreqWavelen), - frequencies * factor, - frequencies - ) - - let isMediumFreq = MLX.logicalAnd( - wavelens .> MLXArray(highFreqWavelen), - wavelens .< MLXArray(lowFreqWavelen) - ) - - let smoothFactors = - (oldContextLen / wavelens - lowFreqFactor) / (highFreqFactor - lowFreqFactor) - let smoothFreqs = frequencies / ((1 - smoothFactors) / factor + smoothFactors) - - freqs = MLX.where(isMediumFreq, smoothFreqs, frequencies) - super.init() - } - - public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { - MLXFast.RoPE( - x, - dimensions: dims, - traditional: traditional, - base: nil, - scale: 1.0, - offset: offset, - freqs: freqs - ) - } - - public func apply(_ x: MLXArray, offset: Int) -> MLXArray { - callAsFunction(x, offset: offset) - } -} - -// MARK: - YarnRoPE - -public class YarnRoPE: Module, RoPEProvider { - let dimensions: Int - let traditional: Bool - let maxPositionEmbeddings: Int - let base: Float - let scalingFactor: Float - let originalMaxPositionEmbeddings: Int - let betaFast: Float - let betaSlow: Float - let mscale: Float - let mscaleAllDim: Float - - private let _mscale: Float - private let _freqs: MLXArray - - public init( - dimensions: Int, - traditional: Bool = false, - maxPositionEmbeddings: Int = 2048, - base: Float = 10000, - scalingFactor: Float = 1.0, - originalMaxPositionEmbeddings: Int = 4096, - betaFast: Float = 32, - betaSlow: Float = 1, - mscale: Float = 1, - mscaleAllDim: Float = 0 - ) { - precondition(dimensions % 2 == 0, "Dimensions must be even") - - self.dimensions = dimensions - self.traditional = traditional - self.maxPositionEmbeddings = maxPositionEmbeddings - self.base = base - self.scalingFactor = scalingFactor - self.originalMaxPositionEmbeddings = originalMaxPositionEmbeddings - self.betaFast = betaFast - self.betaSlow = betaSlow - self.mscale = mscale - self.mscaleAllDim = mscaleAllDim - - func yarnFindCorrectionDim(numRotations: Float) -> Float { - Float(dimensions) - * log(Float(originalMaxPositionEmbeddings) / (numRotations * 2 * Float.pi)) - / (2 * log(base)) - } - - func yarnFindCorrectionRange() -> (low: Int, high: Int) { - let low = Int(floor(yarnFindCorrectionDim(numRotations: betaFast))) - let high = Int(ceil(yarnFindCorrectionDim(numRotations: betaSlow))) - return (max(low, 0), min(high, dimensions - 1)) - } - - func yarnGetMscale(scale: Float, mscale: Float) -> Float { - if scale <= 1 { - return 1.0 - } - return 0.1 * mscale * log(scale) + 1.0 - } - - func yarnLinearRampMask(minVal: Float, maxVal: Float, dim: Int) -> MLXArray { - var maxVal = maxVal - if minVal == maxVal { - maxVal += 0.001 - } - - let linearFunc = (MLXArray(0 ..< dim).asType(.float32) - minVal) / (maxVal - minVal) - return clip(linearFunc, min: 0, max: 1) - } - - _mscale = - yarnGetMscale(scale: scalingFactor, mscale: mscale) - / yarnGetMscale(scale: scalingFactor, mscale: mscaleAllDim) - - let freqExtra = pow( - base, - MLXArray(stride(from: 0, to: dimensions, by: 2)).asType(.float32) - / dimensions - ) - let freqInter = - scalingFactor - * pow( - base, - MLXArray(stride(from: 0, to: dimensions, by: 2)).asType(.float32) - / dimensions - ) - - let (low, high) = yarnFindCorrectionRange() - let freqMask = - 1.0 - yarnLinearRampMask(minVal: Float(low), maxVal: Float(high), dim: dimensions / 2) - - _freqs = (freqInter * freqExtra) / (freqInter * freqMask + freqExtra * (1 - freqMask)) - super.init() - } - - public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { - let input: MLXArray - if _mscale != 1.0 { - // MLXArray subscript assignment creates a new array - x[.ellipsis, 0 ..< dimensions] = _mscale * x[.ellipsis, 0 ..< dimensions] - input = x - } else { - input = x - } - - return MLXFast.RoPE( - input, - dimensions: dimensions, - traditional: traditional, - base: nil, - scale: 1.0, - offset: offset, - freqs: _freqs - ) - } - - public func apply(_ x: MLXArray, offset: Int) -> MLXArray { - callAsFunction(x, offset: offset) - } -} - -// MARK: - SuScaledRoPE (for longrope) - -public class SuScaledRoPE: Module, RoPEProvider { - let dimensions: Int - let base: Float - let maxPositionEmbeddings: Int - let originalMaxPositionEmbeddings: Int - let shortFactor: [Float] - let longFactor: [Float] - - private let shortFreqs: MLXArray - private let longFreqs: MLXArray - private let mscaleShort: Float - private let mscaleLong: Float - - public init( - dimensions: Int, - base: Float = 10000, - maxPositionEmbeddings: Int = 131_072, - originalMaxPositionEmbeddings: Int = 4096, - shortFactor: [Float], - longFactor: [Float] - ) { - self.dimensions = dimensions - self.base = base - self.maxPositionEmbeddings = maxPositionEmbeddings - self.originalMaxPositionEmbeddings = originalMaxPositionEmbeddings - self.shortFactor = shortFactor - self.longFactor = longFactor - - // Compute base frequencies - let baseFreqs = pow( - base, - MLXArray(stride(from: 0, to: dimensions, by: 2)).asType(.float32) - / Float(dimensions) - ) - - // Scale frequencies - shortFreqs = baseFreqs / MLXArray(shortFactor).asType(.float32) - longFreqs = baseFreqs / MLXArray(longFactor).asType(.float32) - - // Compute mscale - let scale = Float(maxPositionEmbeddings) / Float(originalMaxPositionEmbeddings) - if scale <= 1.0 { - mscaleShort = 1.0 - mscaleLong = 1.0 - } else { - mscaleShort = sqrt(1 + log(scale) / log(Float(originalMaxPositionEmbeddings))) - mscaleLong = sqrt(1 + log(scale) / log(Float(originalMaxPositionEmbeddings))) - } - - super.init() - } - - public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { - // Use long freqs when context exceeds original max - let seqLen = x.dim(2) + offset - let freqs = seqLen > originalMaxPositionEmbeddings ? longFreqs : shortFreqs - let mscale = seqLen > originalMaxPositionEmbeddings ? mscaleLong : mscaleShort - - var xMut = x - if mscale != 1.0 { - xMut = mscale * xMut - } - - return MLXFast.RoPE( - xMut, - dimensions: dimensions, - traditional: false, - base: nil, - scale: 1.0, - offset: offset, - freqs: freqs - ) - } - - public func apply(_ x: MLXArray, offset: Int) -> MLXArray { - callAsFunction(x, offset: offset) - } -} - -// MARK: - RoPE Factory - -/// Initialize the appropriate RoPE module based on config -public func initializeRope( - dims: Int, - base: Float, - traditional: Bool, - scalingConfig: [String: StringOrNumber]?, - maxPositionEmbeddings: Int? -) -> any RoPEProvider { - let ropeType: String = { - if let config = scalingConfig, - let typeValue = config["type"] ?? config["rope_type"], - case let .string(s) = typeValue - { - return s - } - return "default" - }() - - if ropeType == "default" || ropeType == "linear" { - let scale: Float = if ropeType == "linear", let factor = scalingConfig?["factor"]?.asFloat() { - 1 / factor - } else { - 1.0 - } - return RoPE(dimensions: dims, traditional: traditional, base: base, scale: scale) - } else if ropeType == "llama3" { - return Llama3RoPE( - dims: dims, - maxPositionEmbeddings: maxPositionEmbeddings ?? 2048, - traditional: traditional, - base: base, - scalingConfig: scalingConfig - ) - } else if ropeType == "yarn" { - let factor = scalingConfig?["factor"]?.asFloat() ?? 32.0 - let origMax = scalingConfig?["original_max_position_embeddings"]?.asInt() ?? 4096 - let betaFast = scalingConfig?["beta_fast"]?.asFloat() ?? 32.0 - let betaSlow = scalingConfig?["beta_slow"]?.asFloat() ?? 1.0 - let mscale = scalingConfig?["mscale"]?.asFloat() ?? 1.0 - let mscaleAllDim = scalingConfig?["mscale_all_dim"]?.asFloat() ?? 0.0 - - return YarnRoPE( - dimensions: dims, - traditional: traditional, - maxPositionEmbeddings: maxPositionEmbeddings ?? 2048, - base: base, - scalingFactor: factor, - originalMaxPositionEmbeddings: origMax, - betaFast: betaFast, - betaSlow: betaSlow, - mscale: mscale, - mscaleAllDim: mscaleAllDim - ) - } else if ropeType == "longrope" { - guard let config = scalingConfig else { - fatalError("longrope requires scaling_config") - } - guard let origMax = config["original_max_position_embeddings"]?.asInt() else { - fatalError("longrope requires original_max_position_embeddings") - } - guard let shortFactor = config["short_factor"]?.asFloats() else { - fatalError("longrope requires short_factor") - } - guard let longFactor = config["long_factor"]?.asFloats() else { - fatalError("longrope requires long_factor") - } - - return SuScaledRoPE( - dimensions: dims, - base: base, - maxPositionEmbeddings: maxPositionEmbeddings ?? 131_072, - originalMaxPositionEmbeddings: origMax, - shortFactor: shortFactor, - longFactor: longFactor - ) - } else if ropeType == "mrope" { - // MRoPE returns basic RoPE here. The actual multi-modal rotary embedding logic - // is handled in the attention layer of multimodal models. - return RoPE(dimensions: dims, traditional: traditional, base: base, scale: 1.0) - } else { - fatalError("Unsupported RoPE type: \(ropeType)") - } -} diff --git a/packages/swift/Sources/NodeMLXCore/StringOrNumber.swift b/packages/swift/Sources/NodeMLXCore/StringOrNumber.swift index 5b5e6ce..b151c49 100644 --- a/packages/swift/Sources/NodeMLXCore/StringOrNumber.swift +++ b/packages/swift/Sources/NodeMLXCore/StringOrNumber.swift @@ -1,105 +1,104 @@ -// Copyright © 2024 Apple Inc. -// Adapted for NodeMLXCore +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Helper type for decoding JSON values that can be either string or number. import Foundation -/// Representation of a heterogenous type in a JSON configuration file. +/// A type that can decode either a string or a number from JSON. /// -/// This can be: a string, a numeric value or an array of numeric values. -/// There are methods to do unwrapping, see e.g. ``asFloat()`` and -/// ``asFloats()`` or callers can switch on the enum. -public enum StringOrNumber: Codable, Equatable, Sendable { +/// This is commonly needed for HuggingFace model configs where some +/// fields may be specified as either strings or numbers (e.g., rope_scaling). +public enum StringOrNumber: Codable, Hashable, Sendable { case string(String) case int(Int) - case float(Float) - case ints([Int]) - case floats([Float]) - case bool(Bool) + case double(Double) public init(from decoder: Decoder) throws { - let values = try decoder.singleValueContainer() + let container = try decoder.singleValueContainer() - if let v = try? values.decode(Int.self) { - self = .int(v) - } else if let v = try? values.decode(Float.self) { - self = .float(v) - } else if let v = try? values.decode([Int].self) { - self = .ints(v) - } else if let v = try? values.decode([Float].self) { - self = .floats(v) - } else if let v = try? values.decode(Bool.self) { - self = .bool(v) + if let intValue = try? container.decode(Int.self) { + self = .int(intValue) + } else if let doubleValue = try? container.decode(Double.self) { + self = .double(doubleValue) + } else if let stringValue = try? container.decode(String.self) { + self = .string(stringValue) } else { - let v = try values.decode(String.self) - self = .string(v) + throw DecodingError.typeMismatch( + StringOrNumber.self, + DecodingError.Context( + codingPath: decoder.codingPath, + debugDescription: "Expected String, Int, or Double" + ) + ) } } public func encode(to encoder: Encoder) throws { var container = encoder.singleValueContainer() switch self { - case let .string(v): try container.encode(v) - case let .int(v): try container.encode(v) - case let .float(v): try container.encode(v) - case let .ints(v): try container.encode(v) - case let .floats(v): try container.encode(v) - case let .bool(v): try container.encode(v) + case let .string(value): + try container.encode(value) + case let .int(value): + try container.encode(value) + case let .double(value): + try container.encode(value) } } - /// Return the value as an optional array of integers. - /// - /// This will not coerce `Float` or `String` to `Int`. - public func asInts() -> [Int]? { + /// Returns the value as a String, converting numbers if necessary. + public var stringValue: String { switch self { - case .string: nil - case let .int(v): [v] - case .float: nil - case let .ints(array): array - case .floats: nil - case .bool: nil + case let .string(value): value + case let .int(value): String(value) + case let .double(value): String(value) } } - /// Return the value as an optional integer. - /// - /// This will not coerce `Float` or `String` to `Int`. - public func asInt() -> Int? { + /// Returns the value as an Int if possible. + public var intValue: Int? { switch self { case .string: nil - case let .int(v): v - case .float: nil - case let .ints(array): array.count == 1 ? array[0] : nil - case .floats: nil - case let .bool(bool): bool ? 1 : 0 + case let .int(value): value + case let .double(value): Int(value) } } - /// Return the value as an optional array of floats. - /// - /// This will not coerce `Int` or `String` to `Float`. - public func asFloats() -> [Float]? { + /// Returns the value as a Double if possible. + public var doubleValue: Double? { switch self { case .string: nil - case let .int(v): [Float(v)] - case let .float(float): [float] - case let .ints(array): array.map { Float($0) } - case let .floats(array): array - case let .bool(bool): [bool ? 1.0 : 0.0] + case let .int(value): Double(value) + case let .double(value): value } } - /// Return the value as an optional float. - /// - /// This will not coerce `Int` or `String` to `Float`. - public func asFloat() -> Float? { + /// Returns the value as a Float if possible. + public var floatValue: Float? { switch self { case .string: nil - case let .int(v): Float(v) - case let .float(float): float - case let .ints(array): array.count == 1 ? Float(array[0]) : nil - case let .floats(array): array.count == 1 ? array[0] : nil - case let .bool(bool): bool ? 1.0 : 0.0 + case let .int(value): Float(value) + case let .double(value): Float(value) + } + } +} + +// MARK: - Dictionary Convenience + +public extension [String: StringOrNumber] { + /// Converts the dictionary to a standard [String: Any] dictionary. + var asAnyDict: [String: Any] { + var result: [String: Any] = [:] + for (key, value) in self { + switch value { + case let .string(s): + result[key] = s + case let .int(i): + result[key] = i + case let .double(d): + result[key] = d + } } + return result } } diff --git a/packages/swift/Sources/NodeMLXCore/Tokenizer.swift b/packages/swift/Sources/NodeMLXCore/Tokenizer.swift index ed6ace0..0b2c6a2 100644 --- a/packages/swift/Sources/NodeMLXCore/Tokenizer.swift +++ b/packages/swift/Sources/NodeMLXCore/Tokenizer.swift @@ -1,191 +1,151 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// Tokenizer.swift -// NodeMLXCore -// -// Tokenizer wrapper using HuggingFace swift-transformers. -// swift-transformers is Apache 2.0 licensed by HuggingFace. -// See: https://github.com/huggingface/swift-transformers -// +// Tokenizer wrapper for HuggingFace tokenizers via swift-transformers. import Foundation import Hub +import MLX import Tokenizers // MARK: - Tokenizer Protocol -/// Protocol for tokenizers that can encode and decode text -public protocol TokenizerProtocol: Sendable { - /// Encode text to token IDs - func encode(_ text: String) -> [Int] +/// Protocol for text tokenization. +public protocol TokenizerProtocol { + /// Encodes text into token IDs. + func encode(text: String) -> [Int] - /// Decode token IDs to text - func decode(_ tokens: [Int]) -> String + /// Decodes token IDs back to text. + func decode(tokens: [Int]) -> String - /// Get special token IDs - var bosTokenId: Int? { get } + /// The vocabulary size. + var vocabularySize: Int { get } + + /// End of sequence token ID. var eosTokenId: Int? { get } - var padTokenId: Int? { get } + + /// Beginning of sequence token ID. + var bosTokenId: Int? { get } } -// MARK: - HuggingFace Tokenizer Wrapper +// MARK: - HuggingFace Tokenizer -/// Wrapper around HuggingFace's Tokenizer from swift-transformers -public class HFTokenizer: TokenizerProtocol, @unchecked Sendable { - private let tokenizer: any Tokenizer +/// Tokenizer loaded from HuggingFace Hub. +public class HFTokenizer: TokenizerProtocol { + private let tokenizer: Tokenizer + private let config: TokenizerConfig? - public let bosTokenId: Int? + public let vocabularySize: Int public let eosTokenId: Int? - public let padTokenId: Int? - - /// Load tokenizer from a HuggingFace model directory - public init(modelDirectory: URL) async throws { - // Load using swift-transformers AutoTokenizer - tokenizer = try await AutoTokenizer.from(modelFolder: modelDirectory) - - // Extract special tokens - try tokenizer_config.json first, then fallback to config.json - var bos: Int? = nil - var eos: Int? = nil - var pad: Int? = nil - - // Helper to extract token ID (handles both Int and [Int] formats) - func extractTokenId(_ value: Any?) -> Int? { - if let intVal = value as? Int { return intVal } - if let array = value as? [Int], let first = array.first { return first } - return nil - } + public let bosTokenId: Int? - // Try tokenizer_config.json - let tokenizerConfigURL = modelDirectory.appendingPathComponent("tokenizer_config.json") - if let data = try? Data(contentsOf: tokenizerConfigURL), - let config = try? JSONSerialization.jsonObject(with: data) as? [String: Any] - { - bos = extractTokenId(config["bos_token_id"]) - eos = extractTokenId(config["eos_token_id"]) - pad = extractTokenId(config["pad_token_id"]) + /// Loads a tokenizer from a local directory (async version). + /// + /// - Parameter path: Path to directory containing tokenizer.json + /// - Throws: Error if tokenizer files cannot be loaded + public init(path: String) async throws { + let url = URL(fileURLWithPath: path) + tokenizer = try await AutoTokenizer.from(modelFolder: url) + + // Try to load tokenizer_config.json for special tokens + let configPath = url.appendingPathComponent("tokenizer_config.json") + if let data = try? Data(contentsOf: configPath) { + config = try? JSONDecoder().decode(TokenizerConfig.self, from: data) + } else { + config = nil } - // Fallback to config.json (model config) for any missing values - let modelConfigURL = modelDirectory.appendingPathComponent("config.json") - if let data = try? Data(contentsOf: modelConfigURL), - let config = try? JSONSerialization.jsonObject(with: data) as? [String: Any] - { - if bos == nil { bos = extractTokenId(config["bos_token_id"]) } - if eos == nil { eos = extractTokenId(config["eos_token_id"]) } - if pad == nil { pad = extractTokenId(config["pad_token_id"]) } - } + // Get vocabulary size - use a reasonable default if not available + vocabularySize = 128_000 // Common default for modern models - bosTokenId = bos - eosTokenId = eos - padTokenId = pad + // Extract special token IDs + eosTokenId = config?.eosTokenId ?? tokenizer.eosTokenId + bosTokenId = config?.bosTokenId ?? tokenizer.bosTokenId } - /// Load tokenizer from HuggingFace Hub model ID - public convenience init(modelId: String) async throws { - // Use Hub to get model directory - let hub = HubApi() - let repo = Hub.Repo(id: modelId) - - // Download tokenizer files - let filePatterns = ["tokenizer.json", "tokenizer_config.json", "vocab.*", "merges.txt"] - let modelDir = try await hub.snapshot(from: repo, matching: filePatterns) - - try await self.init(modelDirectory: modelDir) - } - - // MARK: - TokenizerProtocol - - public func encode(_ text: String) -> [Int] { + public func encode(text: String) -> [Int] { tokenizer.encode(text: text) } - public func decode(_ tokens: [Int]) -> String { + public func decode(tokens: [Int]) -> String { tokenizer.decode(tokens: tokens) } } -// MARK: - Convenience Extensions +// MARK: - Tokenizer Config -public extension HFTokenizer { - /// Encode text with special tokens (BOS/EOS) - func encodeWithSpecialTokens( - _ text: String, - addBos: Bool = true, - addEos: Bool = false - ) -> [Int] { - var tokens = encode(text) +/// Configuration for tokenizer special tokens. +private struct TokenizerConfig: Decodable { + let eosTokenId: Int? + let bosTokenId: Int? + let padTokenId: Int? - if addBos, let bos = bosTokenId { - tokens.insert(bos, at: 0) - } - - if addEos, let eos = eosTokenId { - tokens.append(eos) - } - - return tokens + enum CodingKeys: String, CodingKey { + case eosTokenId = "eos_token_id" + case bosTokenId = "bos_token_id" + case padTokenId = "pad_token_id" } - /// Decode tokens, optionally skipping special tokens - func decode(_ tokens: [Int], skipSpecialTokens: Bool) -> String { - tokenizer.decode(tokens: tokens, skipSpecialTokens: skipSpecialTokens) - } + init(from decoder: Swift.Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) - /// Apply chat template to format a user message for the model - /// Returns token IDs ready for model input - func applyChatTemplate(userMessage: String) throws -> [Int] { - let messages: [[String: any Sendable]] = [ - ["role": "user", "content": userMessage], - ] - return try tokenizer.applyChatTemplate(messages: messages) - } + // Handle both single int and array format + if let id = try? container.decode(Int.self, forKey: .eosTokenId) { + eosTokenId = id + } else if let ids = try? container.decode([Int].self, forKey: .eosTokenId), let first = ids.first { + eosTokenId = first + } else { + eosTokenId = nil + } - /// Apply chat template with conversation history - func applyChatTemplate(messages: [[String: any Sendable]]) throws -> [Int] { - try tokenizer.applyChatTemplate(messages: messages) + if let id = try? container.decode(Int.self, forKey: .bosTokenId) { + bosTokenId = id + } else if let ids = try? container.decode([Int].self, forKey: .bosTokenId), let first = ids.first { + bosTokenId = first + } else { + bosTokenId = nil + } + + if let id = try? container.decode(Int.self, forKey: .padTokenId) { + padTokenId = id + } else if let ids = try? container.decode([Int].self, forKey: .padTokenId), let first = ids.first { + padTokenId = first + } else { + padTokenId = nil + } } } -// MARK: - HuggingFace Hub Utilities - -/// Simple HuggingFace Hub cache utilities -public enum HFHubCache { - public static let cacheDir: URL = { - let home = FileManager.default.homeDirectoryForCurrentUser - return home.appendingPathComponent(".cache/huggingface/hub") - }() - - /// Get the local cache path for a HuggingFace model - public static func modelPath(for modelId: String) -> URL { - let sanitized = modelId.replacingOccurrences(of: "/", with: "--") - return cacheDir - .appendingPathComponent("models--\(sanitized)") - .appendingPathComponent("snapshots") - } +// MARK: - Chat Template - /// Check if a model is cached locally - public static func isCached(_ modelId: String) -> Bool { - let path = modelPath(for: modelId) - var isDir: ObjCBool = false - return FileManager.default.fileExists(atPath: path.path, isDirectory: &isDir) && isDir.boolValue - } +/// Applies chat template to messages. +public func applyChatTemplate( + messages: [[String: String]], + addGenerationPrompt: Bool = true +) -> String { + // Default template for models without explicit chat template + var result = "" - /// Get the latest snapshot directory for a cached model - public static func latestSnapshot(for modelId: String) -> URL? { - let snapshotsDir = modelPath(for: modelId) + for message in messages { + guard let role = message["role"], let content = message["content"] else { + continue + } - guard let contents = try? FileManager.default.contentsOfDirectory( - at: snapshotsDir, - includingPropertiesForKeys: [.contentModificationDateKey], - options: [.skipsHiddenFiles] - ) else { - return nil + switch role { + case "system": + result += "<|system|>\n\(content)\n" + case "user": + result += "<|user|>\n\(content)\n" + case "assistant": + result += "<|assistant|>\n\(content)\n" + default: + result += "\(content)\n" } + } - // Return the most recent snapshot - return contents.sorted { a, b in - let aDate = (try? a.resourceValues(forKeys: [.contentModificationDateKey]).contentModificationDate) ?? .distantPast - let bDate = (try? b.resourceValues(forKeys: [.contentModificationDateKey]).contentModificationDate) ?? .distantPast - return aDate > bDate - }.first + if addGenerationPrompt { + result += "<|assistant|>\n" } + + return result } diff --git a/packages/swift/Sources/NodeMLXCore/Vision/Gemma3VLM.swift b/packages/swift/Sources/NodeMLXCore/Vision/Gemma3VLM.swift deleted file mode 100644 index 50e0ea7..0000000 --- a/packages/swift/Sources/NodeMLXCore/Vision/Gemma3VLM.swift +++ /dev/null @@ -1,312 +0,0 @@ -// -// Gemma3VLM.swift -// NodeMLXCore -// -// Gemma 3 Vision-Language Model -// Combines SigLIP vision encoder with Gemma 3 text model for multimodal generation. -// -// Supports Gemma 3 4B, 12B, and 27B vision variants. -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Configuration - -public struct Gemma3VLMConfiguration: Decodable, Sendable { - /// Text model configuration - public var textConfig: Gemma3Configuration - - /// Vision model configuration - public var visionConfig: SiglipVisionConfiguration - - /// Number of image tokens per image (default: 256) - public var mmTokensPerImage: Int - - /// Begin-of-image token index - public var boiTokenIndex: Int - - /// End-of-image token index - public var eoiTokenIndex: Int - - /// Image placeholder token index - public var imageTokenIndex: Int - - enum CodingKeys: String, CodingKey { - case textConfig = "text_config" - case visionConfig = "vision_config" - case mmTokensPerImage = "mm_tokens_per_image" - case boiTokenIndex = "boi_token_index" - case eoiTokenIndex = "eoi_token_index" - case imageTokenIndex = "image_token_index" - } - - public init(from decoder: Decoder) throws { - let container = try decoder.container(keyedBy: CodingKeys.self) - - textConfig = try container.decode(Gemma3Configuration.self, forKey: .textConfig) - visionConfig = try container.decode(SiglipVisionConfiguration.self, forKey: .visionConfig) - mmTokensPerImage = try container.decodeIfPresent(Int.self, forKey: .mmTokensPerImage) ?? 256 - boiTokenIndex = try container.decodeIfPresent(Int.self, forKey: .boiTokenIndex) ?? 255_999 - eoiTokenIndex = try container.decodeIfPresent(Int.self, forKey: .eoiTokenIndex) ?? 256_000 - imageTokenIndex = try container.decodeIfPresent(Int.self, forKey: .imageTokenIndex) ?? 262_144 - } - - public init( - textConfig: Gemma3Configuration, - visionConfig: SiglipVisionConfiguration, - mmTokensPerImage: Int = 256, - boiTokenIndex: Int = 255_999, - eoiTokenIndex: Int = 256_000, - imageTokenIndex: Int = 262_144 - ) { - self.textConfig = textConfig - self.visionConfig = visionConfig - self.mmTokensPerImage = mmTokensPerImage - self.boiTokenIndex = boiTokenIndex - self.eoiTokenIndex = eoiTokenIndex - self.imageTokenIndex = imageTokenIndex - } -} - -// MARK: - Vision Language Model - -/// Gemma 3 Vision-Language Model -public class Gemma3VLMModel: Module, LLMModel { - // LLMModel protocol - public var vocabularySize: Int { config.textConfig.vocabSize } - public var numLayers: Int { config.textConfig.numHiddenLayers } - public let numKVHeads: Int - public let headDim: Int - public var supportsCache: Bool { true } - - /// Vision tower (SigLIP encoder) - @ModuleInfo(key: "vision_tower") var visionTower: SiglipVisionModel - - /// Multi-modal projector - @ModuleInfo(key: "multi_modal_projector") var multiModalProjector: Gemma3MultiModalProjector - - /// Language model (Gemma 3 text model, wrapped as inner) - @ModuleInfo(key: "language_model") var languageModel: Gemma3Model - - private let config: Gemma3VLMConfiguration - - public init(_ config: Gemma3VLMConfiguration) { - self.config = config - numKVHeads = config.textConfig.numKeyValueHeads - headDim = config.textConfig.headDim - - _visionTower.wrappedValue = SiglipVisionModel(config.visionConfig) - _multiModalProjector.wrappedValue = Gemma3MultiModalProjector( - visionConfig: config.visionConfig, - textHiddenSize: config.textConfig.hiddenSize, - mmTokensPerImage: config.mmTokensPerImage - ) - _languageModel.wrappedValue = Gemma3Model(config.textConfig) - } - - // MARK: - Vision Processing - - /// Get image features from pixel values - /// - Parameter pixelValues: Image tensor [B, C, H, W] - /// - Returns: Projected image features [B, mm_tokens_per_image, hidden_size] - public func getImageFeatures(_ pixelValues: MLXArray) -> MLXArray { - let visionOutputs = visionTower(pixelValues) - let imageFeatures = multiModalProjector(visionOutputs) - return imageFeatures - } - - // MARK: - Forward Pass - - /// Forward pass with optional image input - /// - Parameters: - /// - inputIds: Token IDs [B, L] - /// - pixelValues: Optional image tensor [B, C, H, W] - /// - cache: Optional KV cache - /// - Returns: Logits [B, L, vocab_size] - public func callAsFunction( - _ inputIds: MLXArray, - pixelValues: MLXArray? = nil, - cache: inout [KVCache]? - ) -> MLXArray { - // Get text embeddings - var inputsEmbeds = languageModel.model.embedTokens(inputIds) - - // Scale embeddings (Gemma style) - let scale = MLXArray(sqrt(Float(config.textConfig.hiddenSize))) - inputsEmbeds = inputsEmbeds * scale.asType(inputsEmbeds.dtype) - - // Merge image features if provided - if let pixelValues { - let imageFeatures = getImageFeatures(pixelValues) - inputsEmbeds = mergeImageFeatures(inputsEmbeds, imageFeatures: imageFeatures, inputIds: inputIds) - } - - // Forward through language model with embeddings - return languageModel.forward(inputsEmbeds: inputsEmbeds, cache: &cache) - } - - /// Simple forward without images - public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { - var cache: [KVCache]? = nil - return callAsFunction(inputIds, pixelValues: nil, cache: &cache) - } - - /// Forward with cache (LLMModel protocol) - public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { - callAsFunction(inputIds, pixelValues: nil, cache: &cache) - } - - // MARK: - Cache - - public func newCache() -> [KVCache] { - languageModel.newCache() - } - - // MARK: - Weight Sanitization - - public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - - for (key, value) in weights { - var newKey = key - var newValue = value - - // Map HuggingFace VLM keys to our structure - // HF: vision_tower.vision_model.* -> vision_tower.* (remove redundant vision_model) - if newKey.hasPrefix("vision_tower.vision_model.") { - newKey = "vision_tower." + String(newKey.dropFirst("vision_tower.vision_model.".count)) - } - - // MLX Conv2d expects weights in (out_channels, kH, kW, in_channels) format - // HuggingFace may have (out_channels, in_channels, kH, kW) - need to transpose - if newKey.contains("patch_embedding.weight"), newValue.ndim == 4 { - // Check if format is (out, in, kH, kW) where in=3 for RGB - if newValue.dim(1) == 3, newValue.dim(2) == newValue.dim(3) { - // Transpose from (out, in, kH, kW) to (out, kH, kW, in) - newValue = newValue.transposed(0, 2, 3, 1) - } - } - - result[newKey] = newValue - } - - // Weight tying fallback for language model - if result["language_model.lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["language_model.model.embed_tokens.\(suffix)"] { - result["language_model.lm_head.\(suffix)"] = embedWeight - } - } - } - - return result - } - - // MARK: - Image Feature Merging - - /// Merge image features into text embeddings at placeholder positions - private func mergeImageFeatures( - _ inputsEmbeds: MLXArray, - imageFeatures: MLXArray, - inputIds: MLXArray - ) -> MLXArray { - let imageTokenId = config.imageTokenIndex - let vocabSize = vocabularySize - - // If image token is OOV, handle gracefully - if imageTokenId >= vocabSize { - // Find placeholder positions and replace with image features - // For simplicity, assume single image and replace first occurrence - return maskedScatter(inputsEmbeds, imageFeatures: imageFeatures, inputIds: inputIds, imageTokenId: imageTokenId) - } - - return maskedScatter(inputsEmbeds, imageFeatures: imageFeatures, inputIds: inputIds, imageTokenId: imageTokenId) - } - - /// Scatter image features at mask positions - /// Replaces the single image token with 256 image feature embeddings - private func maskedScatter( - _ inputsEmbeds: MLXArray, - imageFeatures: MLXArray, - inputIds: MLXArray, - imageTokenId: Int - ) -> MLXArray { - let seqLen = inputsEmbeds.dim(1) - - // Find position of image token - var imagePos = -1 - for i in 0 ..< seqLen { - let tokenId = inputIds[0, i].item(Int32.self) - if tokenId == Int32(imageTokenId) { - imagePos = i - break - } - } - - // If no image token found, return unchanged - guard imagePos >= 0 else { - return inputsEmbeds - } - - // Build new embeddings: [before_image] + [image_features] + [after_image] - // This replaces the single image token with 256 image feature tokens - let beforeImage = inputsEmbeds[0..., 0 ..< imagePos, 0...] - let afterImage = inputsEmbeds[0..., (imagePos + 1)..., 0...] - - // Concatenate: before + image_features + after - return concatenated([beforeImage, imageFeatures, afterImage], axis: 1) - } -} - -// MARK: - Gemma3Model Extension for Embeddings Forward - -public extension Gemma3Model { - /// Forward pass with pre-computed embeddings - func forward(inputsEmbeds: MLXArray, cache: inout [KVCache]?) -> MLXArray { - var layerCaches: [KVCache?] = if let existingCache = cache { - existingCache.map { $0 as KVCache? } - } else { - Array(repeating: nil, count: numLayers) - } - - let h = model.forward(inputsEmbeds: inputsEmbeds, cache: &layerCaches) - - cache = layerCaches.compactMap(\.self) - - return lmHead(h) - } -} - -// MARK: - Gemma3ModelInner Extension for Embeddings Forward - -extension Gemma3ModelInner { - /// Forward pass with pre-computed embeddings (skips embed_tokens) - func forward(inputsEmbeds: MLXArray, cache: inout [KVCache?]) -> MLXArray { - var hiddenStates = inputsEmbeds - // Note: embedding scaling should be done before calling this - - // Create masks - let globalLayerIdx = slidingWindowPattern - 1 - let globalCache = globalLayerIdx < cache.count ? cache[globalLayerIdx] : nil - let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) - - let slidingMask: MLXFast.ScaledDotProductAttentionMaskMode - if slidingWindowPattern > 1 { - let firstCache = cache.first ?? nil - slidingMask = createAttentionMask(h: hiddenStates, cache: firstCache, windowSize: slidingWindow) - } else { - slidingMask = globalMask - } - - for i in 0 ..< layers.count { - let isGlobal = (i % slidingWindowPattern) == (slidingWindowPattern - 1) - let mask = isGlobal ? globalMask : slidingMask - hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) - } - - return norm(hiddenStates) - } -} diff --git a/packages/swift/Sources/NodeMLXCore/Vision/ImageProcessor.swift b/packages/swift/Sources/NodeMLXCore/Vision/ImageProcessor.swift deleted file mode 100644 index 4795cc9..0000000 --- a/packages/swift/Sources/NodeMLXCore/Vision/ImageProcessor.swift +++ /dev/null @@ -1,253 +0,0 @@ -// -// ImageProcessor.swift -// NodeMLXCore -// -// Image preprocessing for vision models. -// Handles loading, resizing, and normalizing images. -// - -import CoreGraphics -import Foundation -import ImageIO -import MLX - -// MARK: - Image Processor Configuration - -public struct ImageProcessorConfig: Sendable { - /// Target image size (square) - public var imageSize: Int - - /// Whether to rescale pixel values from [0, 255] to [0, 1] - public var doRescale: Bool - - /// Rescale factor (typically 1/255) - public var rescaleFactor: Float - - /// Whether to normalize with mean/std - public var doNormalize: Bool - - /// Mean values for normalization (per channel RGB) - public var imageMean: [Float] - - /// Std values for normalization (per channel RGB) - public var imageStd: [Float] - - /// Default config for SigLIP/Gemma 3 - public static let siglip = ImageProcessorConfig( - imageSize: 896, - doRescale: true, - rescaleFactor: 1.0 / 255.0, - doNormalize: true, - // SigLIP uses ImageNet normalization - imageMean: [0.5, 0.5, 0.5], - imageStd: [0.5, 0.5, 0.5] - ) - - /// No normalization (just resize) - public static func resizeOnly(size: Int) -> ImageProcessorConfig { - ImageProcessorConfig( - imageSize: size, - doRescale: true, - rescaleFactor: 1.0 / 255.0, - doNormalize: false, - imageMean: [0, 0, 0], - imageStd: [1, 1, 1] - ) - } - - public init( - imageSize: Int = 896, - doRescale: Bool = true, - rescaleFactor: Float = 1.0 / 255.0, - doNormalize: Bool = true, - imageMean: [Float] = [0.5, 0.5, 0.5], - imageStd: [Float] = [0.5, 0.5, 0.5] - ) { - self.imageSize = imageSize - self.doRescale = doRescale - self.rescaleFactor = rescaleFactor - self.doNormalize = doNormalize - self.imageMean = imageMean - self.imageStd = imageStd - } -} - -// MARK: - Image Processor - -/// Preprocesses images for vision models -public struct ImageProcessor: Sendable { - public let config: ImageProcessorConfig - - public init(config: ImageProcessorConfig = .siglip) { - self.config = config - } - - /// Load and preprocess an image from a file path - /// - Parameter path: Path to image file - /// - Returns: Preprocessed image tensor [1, C, H, W] - public func loadAndPreprocess(path: String) throws -> MLXArray { - let url = URL(fileURLWithPath: path) - return try loadAndPreprocess(url: url) - } - - /// Load and preprocess an image from a URL - /// - Parameter url: URL to image file - /// - Returns: Preprocessed image tensor [1, C, H, W] - public func loadAndPreprocess(url: URL) throws -> MLXArray { - let data = try Data(contentsOf: url) - return try preprocess(imageData: data) - } - - /// Preprocess raw image data - /// - Parameter imageData: Raw image bytes (JPEG, PNG, etc.) - /// - Returns: Preprocessed image tensor [1, C, H, W] - public func preprocess(imageData: Data) throws -> MLXArray { - // Decode image using CoreGraphics - guard let provider = CGDataProvider(data: imageData as CFData), - let cgImage = CGImage( - jpegDataProviderSource: provider, - decode: nil, - shouldInterpolate: true, - intent: .defaultIntent - ) ?? CGImage( - pngDataProviderSource: provider, - decode: nil, - shouldInterpolate: true, - intent: .defaultIntent - ) - else { - throw ImageProcessorError.decodeFailed - } - - return preprocess(cgImage: cgImage) - } - - /// Preprocess a CGImage - /// - Parameter cgImage: Core Graphics image - /// - Returns: Preprocessed image tensor [1, C, H, W] - public func preprocess(cgImage: CGImage) -> MLXArray { - // Resize image to target size - let resized = resize(cgImage, to: config.imageSize) - - // Convert to MLXArray [H, W, C] - var pixelValues = cgImageToMLXArray(resized) - - // Rescale from [0, 255] to [0, 1] - if config.doRescale { - pixelValues = pixelValues * config.rescaleFactor - } - - // Normalize with mean/std - if config.doNormalize { - let mean = MLXArray(config.imageMean).reshaped([1, 1, 3]) - let std = MLXArray(config.imageStd).reshaped([1, 1, 3]) - pixelValues = (pixelValues - mean) / std - } - - // Convert from [H, W, C] to [1, C, H, W] (NCHW format) - pixelValues = pixelValues.transposed(2, 0, 1) // [C, H, W] - pixelValues = pixelValues.expandedDimensions(axis: 0) // [1, C, H, W] - - return pixelValues.asType(.float32) - } - - /// Preprocess multiple images - /// - Parameter cgImages: Array of Core Graphics images - /// - Returns: Batched preprocessed tensor [B, C, H, W] - public func preprocess(cgImages: [CGImage]) -> MLXArray { - let processed = cgImages.map { preprocess(cgImage: $0) } - return concatenated(processed, axis: 0) - } -} - -// MARK: - Errors - -public enum ImageProcessorError: Error, LocalizedError { - case decodeFailed - case resizeFailed - case invalidFormat - - public var errorDescription: String? { - switch self { - case .decodeFailed: - "Failed to decode image data" - case .resizeFailed: - "Failed to resize image" - case .invalidFormat: - "Invalid image format" - } - } -} - -// MARK: - Helper Functions - -/// Resize a CGImage to target size (square, center crop) -private func resize(_ image: CGImage, to size: Int) -> CGImage { - let width = image.width - let height = image.height - - // Determine crop region (center crop to square) - let minDim = min(width, height) - let cropX = (width - minDim) / 2 - let cropY = (height - minDim) / 2 - let cropRect = CGRect(x: cropX, y: cropY, width: minDim, height: minDim) - - // Crop to square - guard let croppedImage = image.cropping(to: cropRect) else { - return image - } - - // Create context for resized image - let colorSpace = CGColorSpaceCreateDeviceRGB() - guard let context = CGContext( - data: nil, - width: size, - height: size, - bitsPerComponent: 8, - bytesPerRow: size * 4, - space: colorSpace, - bitmapInfo: CGImageAlphaInfo.noneSkipLast.rawValue - ) else { - return croppedImage - } - - // Draw resized image - context.interpolationQuality = .high - context.draw(croppedImage, in: CGRect(x: 0, y: 0, width: size, height: size)) - - return context.makeImage() ?? croppedImage -} - -/// Convert CGImage to MLXArray [H, W, C] -private func cgImageToMLXArray(_ image: CGImage) -> MLXArray { - let width = image.width - let height = image.height - - // Create RGBA context - let colorSpace = CGColorSpaceCreateDeviceRGB() - var pixelData = [UInt8](repeating: 0, count: width * height * 4) - - guard let context = CGContext( - data: &pixelData, - width: width, - height: height, - bitsPerComponent: 8, - bytesPerRow: width * 4, - space: colorSpace, - bitmapInfo: CGImageAlphaInfo.noneSkipLast.rawValue - ) else { - return MLXArray.zeros([height, width, 3]) - } - - context.draw(image, in: CGRect(x: 0, y: 0, width: width, height: height)) - - // Extract RGB channels (skip alpha) - var rgbData = [Float](repeating: 0, count: width * height * 3) - for i in 0 ..< (width * height) { - rgbData[i * 3 + 0] = Float(pixelData[i * 4 + 0]) // R - rgbData[i * 3 + 1] = Float(pixelData[i * 4 + 1]) // G - rgbData[i * 3 + 2] = Float(pixelData[i * 4 + 2]) // B - } - - return MLXArray(rgbData).reshaped([height, width, 3]) -} diff --git a/packages/swift/Sources/NodeMLXCore/Vision/MultiModalProjector.swift b/packages/swift/Sources/NodeMLXCore/Vision/MultiModalProjector.swift deleted file mode 100644 index f33891d..0000000 --- a/packages/swift/Sources/NodeMLXCore/Vision/MultiModalProjector.swift +++ /dev/null @@ -1,144 +0,0 @@ -// -// MultiModalProjector.swift -// NodeMLXCore -// -// Multi-Modal Projector for Gemma 3 VLM. -// Projects vision embeddings into the language model's embedding space. -// -// Based on HuggingFace transformers Gemma3MultiModalProjector -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Gemma3 RMSNorm (for projector) - -/// RMSNorm with Gemma-style (1 + weight) scaling -public class ProjectorRMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - public init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.zeros([dimensions]) - } - - public func callAsFunction(_ x: MLXArray) -> MLXArray { - // Gemma uses (1 + weight) scaling - MLXFast.rmsNorm(x, weight: 1 + weight, eps: eps) - } -} - -// MARK: - Multi-Modal Projector - -/// Projects vision features into language model space -/// Uses average pooling to reduce patch count to mm_tokens_per_image (256 for Gemma 3) -public class Gemma3MultiModalProjector: Module { - /// Linear projection weight (manual parameter, not wrapped Linear) - @ModuleInfo(key: "mm_input_projection_weight") var projectionWeight: MLXArray - - /// RMSNorm before projection - @ModuleInfo(key: "mm_soft_emb_norm") var softEmbNorm: ProjectorRMSNorm - - let patchesPerImage: Int - let tokensPerSide: Int - let kernelSize: Int - - /// Initialize the projector - /// - Parameters: - /// - visionHiddenSize: Hidden size of vision encoder (e.g., 1152) - /// - textHiddenSize: Hidden size of text model (e.g., 2304) - /// - imageSize: Vision model image size (e.g., 896) - /// - patchSize: Vision model patch size (e.g., 14) - /// - mmTokensPerImage: Target number of image tokens (e.g., 256) - /// - layerNormEps: Layer norm epsilon - public init( - visionHiddenSize: Int, - textHiddenSize: Int, - imageSize: Int = 896, - patchSize: Int = 14, - mmTokensPerImage: Int = 256, - layerNormEps: Float = 1e-6 - ) { - // Calculate pooling parameters - patchesPerImage = imageSize / patchSize // 896/14 = 64 - tokensPerSide = Int(sqrt(Double(mmTokensPerImage))) // sqrt(256) = 16 - kernelSize = patchesPerImage / tokensPerSide // 64/16 = 4 - - // Initialize projection weight to zeros (following HF init) - _projectionWeight.wrappedValue = MLXArray.zeros([visionHiddenSize, textHiddenSize]) - - // RMSNorm for soft embeddings - _softEmbNorm.wrappedValue = ProjectorRMSNorm(dimensions: visionHiddenSize, eps: layerNormEps) - } - - /// Initialize from configs - public init(visionConfig: SiglipVisionConfiguration, textHiddenSize: Int, mmTokensPerImage: Int = 256) { - patchesPerImage = visionConfig.patchesPerSide - tokensPerSide = Int(sqrt(Double(mmTokensPerImage))) - kernelSize = patchesPerImage / tokensPerSide - - _projectionWeight.wrappedValue = MLXArray.zeros([visionConfig.hiddenSize, textHiddenSize]) - _softEmbNorm.wrappedValue = ProjectorRMSNorm(dimensions: visionConfig.hiddenSize, eps: visionConfig.layerNormEps) - } - - /// Project vision features to language model space - /// - Parameter visionOutputs: Vision encoder output [B, num_patches, vision_hidden] - /// - Returns: Projected features [B, mm_tokens_per_image, text_hidden] - public func callAsFunction(_ visionOutputs: MLXArray) -> MLXArray { - let batchSize = visionOutputs.dim(0) - let seqLength = visionOutputs.dim(2) // vision_hidden_size - - // Reshape for 2D pooling: [B, num_patches, hidden] -> [B, hidden, patches_h, patches_w] - var reshaped = visionOutputs.transposed(0, 2, 1) // [B, hidden, num_patches] - reshaped = reshaped.reshaped([batchSize, seqLength, patchesPerImage, patchesPerImage]) - - // Average pooling to reduce spatial dimensions - // [B, hidden, 64, 64] -> [B, hidden, 16, 16] with kernel_size=4 - let pooled = avgPool2d(reshaped, kernelSize: kernelSize) - - // Flatten spatial dims: [B, hidden, tokens_h, tokens_w] -> [B, hidden, num_tokens] - let flattened = pooled.reshaped([batchSize, seqLength, -1]) - - // Transpose back: [B, hidden, num_tokens] -> [B, num_tokens, hidden] - var output = flattened.transposed(0, 2, 1) - - // Apply RMSNorm - output = softEmbNorm(output) - - // Project to text hidden size: [B, num_tokens, vision_hidden] @ [vision_hidden, text_hidden] - output = matmul(output, projectionWeight) - - return output - } -} - -// MARK: - Average Pooling Helper - -/// 2D Average Pooling -/// - Parameters: -/// - x: Input tensor [B, C, H, W] -/// - kernelSize: Pooling kernel size -/// - Returns: Pooled tensor [B, C, H/kernel, W/kernel] -private func avgPool2d(_ x: MLXArray, kernelSize: Int) -> MLXArray { - let (B, C, H, W) = (x.dim(0), x.dim(1), x.dim(2), x.dim(3)) - let newH = H / kernelSize - let newW = W / kernelSize - - // Reshape to extract pooling windows - // [B, C, H, W] -> [B, C, newH, kernel, newW, kernel] - var reshaped = x.reshaped([B, C, newH, kernelSize, newW, kernelSize]) - - // Move kernel dims together and compute mean - // [B, C, newH, kernel, newW, kernel] -> [B, C, newH, newW, kernel, kernel] - reshaped = reshaped.transposed(0, 1, 2, 4, 3, 5) - - // Reshape to [B, C, newH, newW, kernel*kernel] and mean over last dim - reshaped = reshaped.reshaped([B, C, newH, newW, kernelSize * kernelSize]) - let pooled = reshaped.mean(axis: -1) - - return pooled -} diff --git a/packages/swift/Sources/NodeMLXCore/Vision/SiglipVision.swift b/packages/swift/Sources/NodeMLXCore/Vision/SiglipVision.swift deleted file mode 100644 index b7b33c1..0000000 --- a/packages/swift/Sources/NodeMLXCore/Vision/SiglipVision.swift +++ /dev/null @@ -1,308 +0,0 @@ -// -// SiglipVision.swift -// NodeMLXCore -// -// SigLIP Vision Encoder for Gemma 3 VLM. -// Converts images into visual embeddings. -// -// Based on: -// - HuggingFace transformers (Apache 2.0): models/siglip/modeling_siglip.py -// - mlx-vlm (MIT): models/siglip/siglip.py -// - -import Foundation -import MLX -import MLXFast -import MLXNN - -// MARK: - Configuration - -public struct SiglipVisionConfiguration: Decodable, Sendable { - public var hiddenSize: Int - public var intermediateSize: Int - public var numHiddenLayers: Int - public var numAttentionHeads: Int - public var numChannels: Int - public var imageSize: Int - public var patchSize: Int - public var layerNormEps: Float - public var hiddenAct: String - - enum CodingKeys: String, CodingKey { - case hiddenSize = "hidden_size" - case intermediateSize = "intermediate_size" - case numHiddenLayers = "num_hidden_layers" - case numAttentionHeads = "num_attention_heads" - case numChannels = "num_channels" - case imageSize = "image_size" - case patchSize = "patch_size" - case layerNormEps = "layer_norm_eps" - case hiddenAct = "hidden_act" - } - - public init(from decoder: Decoder) throws { - let container = try decoder.container(keyedBy: CodingKeys.self) - - hiddenSize = try container.decodeIfPresent(Int.self, forKey: .hiddenSize) ?? 1152 - intermediateSize = try container.decodeIfPresent(Int.self, forKey: .intermediateSize) ?? 4304 - numHiddenLayers = try container.decodeIfPresent(Int.self, forKey: .numHiddenLayers) ?? 27 - numAttentionHeads = try container.decodeIfPresent(Int.self, forKey: .numAttentionHeads) ?? 16 - numChannels = try container.decodeIfPresent(Int.self, forKey: .numChannels) ?? 3 - imageSize = try container.decodeIfPresent(Int.self, forKey: .imageSize) ?? 896 - patchSize = try container.decodeIfPresent(Int.self, forKey: .patchSize) ?? 14 - layerNormEps = try container.decodeIfPresent(Float.self, forKey: .layerNormEps) ?? 1e-6 - hiddenAct = try container.decodeIfPresent(String.self, forKey: .hiddenAct) ?? "gelu_pytorch_tanh" - } - - public init( - hiddenSize: Int = 1152, - intermediateSize: Int = 4304, - numHiddenLayers: Int = 27, - numAttentionHeads: Int = 16, - numChannels: Int = 3, - imageSize: Int = 896, - patchSize: Int = 14, - layerNormEps: Float = 1e-6, - hiddenAct: String = "gelu_pytorch_tanh" - ) { - self.hiddenSize = hiddenSize - self.intermediateSize = intermediateSize - self.numHiddenLayers = numHiddenLayers - self.numAttentionHeads = numAttentionHeads - self.numChannels = numChannels - self.imageSize = imageSize - self.patchSize = patchSize - self.layerNormEps = layerNormEps - self.hiddenAct = hiddenAct - } - - /// Number of patches per side - public var patchesPerSide: Int { - imageSize / patchSize - } - - /// Total number of patches - public var numPatches: Int { - patchesPerSide * patchesPerSide - } - - /// Head dimension - public var headDim: Int { - hiddenSize / numAttentionHeads - } -} - -// MARK: - Vision Embeddings - -/// Converts image pixels to patch embeddings with positional encoding -public class SiglipVisionEmbeddings: Module { - @ModuleInfo(key: "patch_embedding") var patchEmbedding: Conv2d - @ModuleInfo(key: "position_embedding") var positionEmbedding: Embedding - - let numPatches: Int - - public init(_ config: SiglipVisionConfiguration) { - numPatches = config.numPatches - - // Conv2d to extract patches: [B, C, H, W] -> [B, hidden, patches, patches] - let patchSize = config.patchSize - _patchEmbedding.wrappedValue = Conv2d( - inputChannels: config.numChannels, - outputChannels: config.hiddenSize, - kernelSize: IntOrPair((patchSize, patchSize)), - stride: IntOrPair((patchSize, patchSize)), - padding: IntOrPair((0, 0)) - ) - - // Learnable position embeddings for each patch - _positionEmbedding.wrappedValue = Embedding( - embeddingCount: numPatches, - dimensions: config.hiddenSize - ) - } - - public func callAsFunction(_ pixelValues: MLXArray) -> MLXArray { - // pixelValues: [B, C, H, W] or [B, H, W, C] - var x = pixelValues - - // MLX Conv2d expects NHWC format - if x.dim(1) == 3, x.dim(2) == x.dim(3) { - // Convert NCHW to NHWC - x = x.transposed(0, 2, 3, 1) - } - - // Apply patch embedding: [B, H, W, C] -> [B, patches_h, patches_w, hidden] - let patchEmbeds = patchEmbedding(x) - - // Flatten patches: [B, ph, pw, hidden] -> [B, num_patches, hidden] - let batchSize = patchEmbeds.dim(0) - let hiddenSize = patchEmbeds.dim(3) - var embeddings = patchEmbeds.reshaped([batchSize, -1, hiddenSize]) - - // Add position embeddings - let positionIds = MLXArray(0 ..< numPatches) - embeddings = embeddings + positionEmbedding(positionIds) - - return embeddings - } -} - -// MARK: - Attention - -/// Multi-head self-attention for vision -public class SiglipAttention: Module { - @ModuleInfo(key: "q_proj") var qProj: Linear - @ModuleInfo(key: "k_proj") var kProj: Linear - @ModuleInfo(key: "v_proj") var vProj: Linear - @ModuleInfo(key: "out_proj") var outProj: Linear - - let numHeads: Int - let headDim: Int - let scale: Float - - public init(_ config: SiglipVisionConfiguration) { - numHeads = config.numAttentionHeads - headDim = config.headDim - scale = pow(Float(headDim), -0.5) - - let hiddenSize = config.hiddenSize - _qProj.wrappedValue = Linear(hiddenSize, hiddenSize) - _kProj.wrappedValue = Linear(hiddenSize, hiddenSize) - _vProj.wrappedValue = Linear(hiddenSize, hiddenSize) - _outProj.wrappedValue = Linear(hiddenSize, hiddenSize) - } - - public func callAsFunction(_ hiddenStates: MLXArray) -> MLXArray { - let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - - // Project to Q, K, V - var queries = qProj(hiddenStates) - var keys = kProj(hiddenStates) - var values = vProj(hiddenStates) - - // Reshape: [B, L, hidden] -> [B, L, heads, headDim] -> [B, heads, L, headDim] - queries = queries.reshaped([B, L, numHeads, headDim]).transposed(0, 2, 1, 3) - keys = keys.reshaped([B, L, numHeads, headDim]).transposed(0, 2, 1, 3) - values = values.reshaped([B, L, numHeads, headDim]).transposed(0, 2, 1, 3) - - // Scaled dot-product attention (no causal mask for vision) - let output = MLXFast.scaledDotProductAttention( - queries: queries, - keys: keys, - values: values, - scale: scale, - mask: .none - ) - - // Reshape back: [B, heads, L, headDim] -> [B, L, hidden] - let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) - - return outProj(outputReshaped) - } -} - -// MARK: - MLP - -/// Feed-forward network with GELU activation -public class SiglipMLP: Module { - @ModuleInfo(key: "fc1") var fc1: Linear - @ModuleInfo(key: "fc2") var fc2: Linear - - let useApproxGelu: Bool - - public init(_ config: SiglipVisionConfiguration) { - _fc1.wrappedValue = Linear(config.hiddenSize, config.intermediateSize) - _fc2.wrappedValue = Linear(config.intermediateSize, config.hiddenSize) - - // Check if using approximate GELU (pytorch_tanh variant) - useApproxGelu = config.hiddenAct.contains("tanh") - } - - public func callAsFunction(_ x: MLXArray) -> MLXArray { - var h = fc1(x) - h = useApproxGelu ? geluApproximate(h) : gelu(h) - return fc2(h) - } -} - -// MARK: - Encoder Layer - -/// Single transformer encoder layer -public class SiglipEncoderLayer: Module { - @ModuleInfo(key: "layer_norm1") var layerNorm1: LayerNorm - @ModuleInfo(key: "self_attn") var selfAttn: SiglipAttention - @ModuleInfo(key: "layer_norm2") var layerNorm2: LayerNorm - @ModuleInfo(key: "mlp") var mlp: SiglipMLP - - public init(_ config: SiglipVisionConfiguration) { - _layerNorm1.wrappedValue = LayerNorm(dimensions: config.hiddenSize, eps: config.layerNormEps) - _selfAttn.wrappedValue = SiglipAttention(config) - _layerNorm2.wrappedValue = LayerNorm(dimensions: config.hiddenSize, eps: config.layerNormEps) - _mlp.wrappedValue = SiglipMLP(config) - } - - public func callAsFunction(_ hiddenStates: MLXArray) -> MLXArray { - // Pre-norm self-attention - var residual = hiddenStates - var h = layerNorm1(hiddenStates) - h = selfAttn(h) - h = residual + h - - // Pre-norm MLP - residual = h - h = layerNorm2(h) - h = mlp(h) - h = residual + h - - return h - } -} - -// MARK: - Vision Encoder - -/// Full SigLIP vision encoder -public class SiglipEncoder: Module { - @ModuleInfo(key: "layers") var layers: [SiglipEncoderLayer] - - public init(_ config: SiglipVisionConfiguration) { - _layers.wrappedValue = (0 ..< config.numHiddenLayers).map { _ in - SiglipEncoderLayer(config) - } - } - - public func callAsFunction(_ hiddenStates: MLXArray) -> MLXArray { - var h = hiddenStates - for layer in layers { - h = layer(h) - } - return h - } -} - -// MARK: - Vision Model - -/// Complete SigLIP Vision Model -public class SiglipVisionModel: Module { - @ModuleInfo(key: "embeddings") var embeddings: SiglipVisionEmbeddings - @ModuleInfo(key: "encoder") var encoder: SiglipEncoder - @ModuleInfo(key: "post_layernorm") var postLayernorm: LayerNorm - - let config: SiglipVisionConfiguration - - public init(_ config: SiglipVisionConfiguration) { - self.config = config - _embeddings.wrappedValue = SiglipVisionEmbeddings(config) - _encoder.wrappedValue = SiglipEncoder(config) - _postLayernorm.wrappedValue = LayerNorm(dimensions: config.hiddenSize, eps: config.layerNormEps) - } - - /// Forward pass - /// - Parameter pixelValues: Image tensor [B, C, H, W] or [B, H, W, C] - /// - Returns: Hidden states [B, num_patches, hidden_size] - public func callAsFunction(_ pixelValues: MLXArray) -> MLXArray { - var hiddenStates = embeddings(pixelValues) - hiddenStates = encoder(hiddenStates) - hiddenStates = postLayernorm(hiddenStates) - return hiddenStates - } -} diff --git a/packages/swift/Sources/NodeMLXCore/generated/README.md b/packages/swift/Sources/NodeMLXCore/generated/README.md new file mode 100644 index 0000000..7d73bdc --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/generated/README.md @@ -0,0 +1,127 @@ +# Generated Code + +⚠️ **DO NOT EDIT FILES IN THIS DIRECTORY MANUALLY** ⚠️ + +All files are auto-generated by `hf2swift` and will be overwritten. + +## Models + +| Model | File | HuggingFace Type | Features | +| -------- | ------------------------- | ---------------- | ---------------------------- | +| Llama | `LlamaGenerated.swift` | `llama` | Standard (shared components) | +| Phi-3 | `Phi3Generated.swift` | `phi3` | Fused QKV | +| Qwen2 | `Qwen2Generated.swift` | `qwen2` | Standard (shared components) | +| Qwen3 | `Qwen3Generated.swift` | `qwen3` | Q/K norms | +| Gemma3 | `Gemma3Generated.swift` | `gemma3` | 4 norms, Gemma RMSNorm | +| Gemma3n | `Gemma3nGenerated.swift` | `gemma3n` | AltUp, Laurel, VLM | +| Mistral | `MistralGenerated.swift` | `mistral` | Sliding window | +| Mistral3 | `Mistral3Generated.swift` | `mistral3` | YaRN RoPE | +| SmolLM3 | `SmolLM3Generated.swift` | `smollm3` | No-RoPE layers | +| GPT-OSS | `GptOSSGenerated.swift` | `gpt_oss` | MoE, attention sinks | + +## Regenerating Models + +### Single Model + +```bash +pnpm hf2swift --model llama --output packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift +``` + +### All Models (Automatic) + +The pre-push hook automatically regenerates all models: + +```bash +git push # Regenerates and validates all models +``` + +### Manual Regeneration + +```bash +cd packages/hf2swift +for model in llama qwen2 qwen3 mistral mistral3 phi3 gemma3 gemma3n smollm3 gpt_oss; do + pnpm tsx src/cli.ts --model $model --output ../swift/Sources/NodeMLXCore/generated/models/Generated.swift +done +``` + +## Generator Source + +The generator is at `packages/hf2swift/`: + +``` +hf2swift/src/generator/ +├── model-defs/ # Model family definitions +│ ├── llama.ts # Llama architectural features +│ ├── qwen.ts # Qwen2, Qwen3 +│ ├── gemma.ts # Gemma3, Gemma3n +│ └── ... +├── components/ # Code generators +│ ├── attention.ts # Attention layer +│ ├── mlp.ts # MLP layer +│ ├── decoder-layer.ts # Decoder layer +│ └── model.ts # Model wrapper +└── features.ts # Feature merging logic +``` + +## How Generation Works + +1. **Feature detection**: Determine architectural features from model type +2. **Component selection**: Choose shared vs. custom implementation per component +3. **Code generation**: Produce Swift code +4. **SwiftFormat**: Apply consistent formatting + +### Simple Models (Llama, Qwen2) + +Use shared components via typealiases: + +```swift +typealias LlamaAttention = StandardAttention +typealias LlamaMLP = StandardMLP +typealias LlamaDecoderLayer = StandardDecoderLayer +``` + +Result: ~195 lines of generated code + +### Complex Models (Gemma3n, GPT-OSS) + +Generate custom implementations: + +```swift +class Gemma3nAttention: Module { + // Full custom implementation with Q/K/V norms, sliding window, etc. +} +``` + +Result: ~700+ lines of generated code + +## Validation + +Generated files are validated on every push: + +1. Pre-push hook regenerates all models +2. Compares against committed versions +3. Fails if any differences detected +4. Ensures generator and generated code stay in sync + +## Adding a New Model + +1. Create `model-defs/.ts`: + + ```typescript + export const myModel: ModelDefinition = { + name: "MyModel", + matches: (t) => t.includes("mymodel"), + architectural: { ...DEFAULT_ARCHITECTURAL, activation: "silu" }, + configDefaults: { ...DEFAULT_CONFIG, ropeTheta: 50000 } + } + ``` + +2. Register in `model-defs/index.ts` + +3. Generate: + + ```bash + pnpm hf2swift --model mymodel --output .../MyModelGenerated.swift + ``` + +4. Add to pre-push hook model list diff --git a/packages/swift/Sources/NodeMLXCore/Models/Gemma3Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Gemma3Generated.swift similarity index 85% rename from packages/swift/Sources/NodeMLXCore/Models/Gemma3Generated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/Gemma3Generated.swift index a3c0eac..9df3b19 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/Gemma3Generated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Gemma3Generated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct Gemma3Configuration: Decodable, Sendable { +public struct Gemma3Configuration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -89,7 +89,7 @@ public struct Gemma3Configuration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.000001) ropeTheta = try decode(.ropeTheta, default: 1_000_000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -105,34 +105,14 @@ public struct Gemma3Configuration: Decodable, Sendable { // MARK: - RMS Norm -/// RMSNorm with Gemma-style (1 + weight) scaling -class Gemma3RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - // Initialize to zeros - will be (1 + weight) in forward - _weight.wrappedValue = MLXArray.zeros([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - // Gemma uses (1 + weight) scaling - MLXFast.rmsNorm(x, weight: 1 + weight, eps: eps) - } -} +/// Uses ported GemmaRMSNorm (1 + weight scaling) +typealias Gemma3RMSNorm = GemmaRMSNorm // MARK: - Utility Functions -/// Clip residual for float16 overflow protection (matching mlx-lm) +/// Clip residual for float16 overflow protection - uses shared implementation private func clipResidual(_ x: MLXArray, _ y: MLXArray) -> MLXArray { - if x.dtype != .float16 { - return x + y - } - let bound = Float16.greatestFiniteMagnitude - let sum = (x.asType(.float32) + y.asType(.float32)) - return clip(sum, min: MLXArray(-Float(bound)), max: MLXArray(Float(bound))).asType(.float16) + MathUtils.clipResidual(x, y) } // MARK: - Attention @@ -193,8 +173,8 @@ class Gemma3Attention: Module { // Apply RoPE with cache offset let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) // Update cache if let c = cache { @@ -303,11 +283,12 @@ class Gemma3ModelInner: Module { hiddenStates = hiddenStates * scale.asType(hiddenStates.dtype) let globalLayerIdx = slidingWindowPattern - 1 let globalCache = globalLayerIdx < cache.count ? cache[globalLayerIdx] : nil - let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) + let globalOffset = globalCache?.offset ?? 0 + let globalMask = createAttentionMask(n: hiddenStates.dim(1), offset: globalOffset, windowSize: nil) let slidingMask: MLXFast.ScaledDotProductAttentionMaskMode if slidingWindowPattern > 1 { - let firstCache = cache.first ?? nil - slidingMask = createAttentionMask(h: hiddenStates, cache: firstCache, windowSize: slidingWindow) + let slidingOffset = cache.first??.offset ?? 0 + slidingMask = createAttentionMask(n: hiddenStates.dim(1), offset: slidingOffset, windowSize: slidingWindow) } else { slidingMask = globalMask } @@ -368,20 +349,7 @@ public class Gemma3Model: Module, LLMModel { } public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + // Uses shared weight sanitization logic + sanitizeWeights(weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/Models/Gemma3nGenerated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Gemma3nGenerated.swift similarity index 82% rename from packages/swift/Sources/NodeMLXCore/Models/Gemma3nGenerated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/Gemma3nGenerated.swift index f903f58..3533df7 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/Gemma3nGenerated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Gemma3nGenerated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct Gemma3nConfiguration: Decodable, Sendable { +public struct Gemma3nConfiguration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -47,6 +47,11 @@ public struct Gemma3nConfiguration: Decodable, Sendable { public var ropeScaling: [String: StringOrNumber]? public var modelType: String? + /// Default intermediate size (first layer) for BaseModelConfiguration conformance + public var intermediateSize: Int { + intermediateSizes.first ?? 16384 + } + /// Get intermediate size for a specific layer public func intermediateSize(forLayer idx: Int) -> Int { if idx < intermediateSizes.count { @@ -140,7 +145,7 @@ public struct Gemma3nConfiguration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.000001) ropeTheta = try decode(.ropeTheta, default: 1_000_000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -195,140 +200,26 @@ class RMSNoScale: Module { } } -/// Standard RMSNorm -class Gemma3nRMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} +/// Uses shared RMSNorm implementation +typealias Gemma3nRMSNorm = RMSNorm // MARK: - Utility Functions // MARK: - AltUp Block -/// Alternating Updates module for efficient sparse computation -class Gemma3nAltUp: Module { - let numInputs: Int - let activeIdx: Int - let hiddenSize: Int - let altupCoefClip: Float? - - @ModuleInfo(key: "correct_output_scale") var correctOutputScale: MLXArray - @ModuleInfo(key: "correction_coefs") var correctionCoefs: Linear - @ModuleInfo(key: "prediction_coefs") var predictionCoefs: Linear - @ModuleInfo(key: "modality_router") var modalityRouter: Linear - @ModuleInfo(key: "router_norm") var routerNorm: Gemma3nRMSNorm - - init(_ config: Gemma3nConfiguration) { - numInputs = config.altupNumInputs - activeIdx = config.altupActiveIdx - hiddenSize = config.hiddenSize - altupCoefClip = config.altupCoefClip - - _correctOutputScale.wrappedValue = MLXArray.zeros([config.hiddenSize]) - _correctionCoefs.wrappedValue = Linear(numInputs, numInputs, bias: false) - _predictionCoefs.wrappedValue = Linear(numInputs, numInputs * numInputs, bias: false) - _modalityRouter.wrappedValue = Linear(config.hiddenSize, numInputs, bias: false) - _routerNorm.wrappedValue = Gemma3nRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func computeRouterModalities(_ x: MLXArray) -> MLXArray { - let routerInputs = routerNorm(x) * pow(Float(hiddenSize), -1.0) - let routed = modalityRouter(routerInputs).asType(.float32) - return tanh(routed) - } - - /// Predict step: modifies input using learned coefficients - /// Input: [numInputs, batch, seq, hidden] -> Output: [numInputs, batch, seq, hidden] - func predict(_ hiddenStates: MLXArray) -> MLXArray { - let modalities = computeRouterModalities(hiddenStates[activeIdx]) +/// AltUpConfiguration conformance for shared AltUpBlock +extension Gemma3nConfiguration: AltUpConfiguration {} - // Compute prediction coefficients with optional clipping - var weight = predictionCoefs.weight.asType(.float32) - if let clipVal = altupCoefClip { - weight = clip(weight, min: -clipVal, max: clipVal) - } - - // Manual linear: modalities @ weight.T - var allCoefs = matmul(modalities.asType(.float32), weight.T) - let shape = modalities.shape - allCoefs = allCoefs.reshaped([shape[0], shape[1], numInputs, numInputs]) - allCoefs = allCoefs.transposed(0, 1, 3, 2) - - // Convert to float32 for better precision - let xUp = hiddenStates.asType(.float32) - let xPermuted = xUp.transposed(1, 2, 3, 0) - var predictions = matmul(xPermuted, allCoefs) - predictions = predictions.transposed(3, 0, 1, 2) - predictions = predictions + xUp - - return predictions.asType(hiddenStates.dtype) - } - - /// Correct step: refines predictions based on activated output - func correct(_ predictions: MLXArray, activated: MLXArray) -> MLXArray { - let modalities = computeRouterModalities(activated) - - // Compute correction coefficients with optional clipping - var weight = correctionCoefs.weight.asType(.float32) - if let clipVal = altupCoefClip { - weight = clip(weight, min: -clipVal, max: clipVal) - } - - // Manual linear + 1.0: modalities @ weight.T + 1.0 - var allCoefs = matmul(modalities.asType(.float32), weight.T) + 1.0 - let activeX = predictions[activeIdx] - let innovation = activated - activeX - - // allCoefs: [batch, seq, numInputs] -> [numInputs, batch, seq] - allCoefs = allCoefs.transposed(2, 0, 1) - - // innovation: [batch, seq, hidden] - // We need to broadcast: [numInputs, batch, seq, 1] * [1, batch, seq, hidden] - let innovationExpanded = innovation.expandedDimensions(axis: 0) - let allCoefsExpanded = allCoefs.expandedDimensions(axis: -1) - let corrected = innovationExpanded * allCoefsExpanded + predictions - - return corrected.asType(activated.dtype) - } - - func scaleCorrectOutput(_ corrected: MLXArray) -> MLXArray { - corrected * correctOutputScale - } -} +/// AltUp block - uses shared implementation +typealias Gemma3nAltUp = AltUpBlock // MARK: - Laurel Block -/// Low-rank residual layer (Learned Augmented Residual) -/// Note: This layer adds the residual internally (returns x + laurel_output) -class Gemma3nLaurelBlock: Module { - @ModuleInfo(key: "linear_left") var linearLeft: Linear - @ModuleInfo(key: "linear_right") var linearRight: Linear - @ModuleInfo(key: "post_laurel_norm") var postLaurelNorm: Gemma3nRMSNorm +/// LaurelConfiguration conformance for shared LaurelBlock +extension Gemma3nConfiguration: LaurelConfiguration {} - init(_ config: Gemma3nConfiguration) { - _linearLeft.wrappedValue = Linear(config.hiddenSize, config.laurelRank, bias: false) - _linearRight.wrappedValue = Linear(config.laurelRank, config.hiddenSize, bias: false) - _postLaurelNorm.wrappedValue = Gemma3nRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - var laurel = linearLeft(x) - laurel = linearRight(laurel) - laurel = postLaurelNorm(laurel) - // Add residual connection - return x + laurel - } -} +/// Laurel block - uses shared implementation +typealias Gemma3nLaurelBlock = LaurelBlock // MARK: - Attention @@ -398,7 +289,7 @@ class Gemma3nAttention: Module { keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) keys = kNorm(keys) keys = keys.transposed(0, 2, 1, 3) - keys = rope.apply(keys, offset: offset) + keys = rope(keys, offset: offset) values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) values = vNorm(values) values = values.transposed(0, 2, 1, 3) @@ -406,7 +297,7 @@ class Gemma3nAttention: Module { (keys, values) = c.update(keys: keys, values: values) } } - queries = rope.apply(queries, offset: offset) + queries = rope(queries, offset: offset) // Attention using MLXFast (handles GQA automatically) let output = MLXFast.scaledDotProductAttention( @@ -449,23 +340,12 @@ class Gemma3nMLP: Module { // Precompute std multiplier for gelu_topk if sparsity > 0 if activationSparsity > 0 { // sqrt(2) * erfinv(2 * sparsity - 1) - stdMultiplier = Float(sqrt(2.0)) * Self.erfinv(2.0 * activationSparsity - 1.0) + stdMultiplier = Float(sqrt(2.0)) * MathUtils.erfinv(2.0 * activationSparsity - 1.0) } else { stdMultiplier = nil } } - /// Approximate inverse error function - private static func erfinv(_ x: Float) -> Float { - let a: Float = 0.147 - let sign: Float = x < 0 ? -1 : 1 - let x2 = x * x - let lnTerm = log(1 - x2) - let term1 = 2 / (Float.pi * a) + lnTerm / 2 - let term2 = lnTerm / a - return sign * sqrt(sqrt(term1 * term1 - term2) - term1) - } - func callAsFunction(_ x: MLXArray) -> MLXArray { let gateOutput = gateProj(x) let activations: MLXArray @@ -702,9 +582,11 @@ class Gemma3nLanguageModel: Module { let h0 = hiddenStates[0] let globalCache = firstFullIdx < cache.count ? cache[firstFullIdx] : nil - let globalMask = createAttentionMask(h: h0, cache: globalCache, windowSize: nil) + let globalOffset = globalCache?.offset ?? 0 + let globalMask = createAttentionMask(n: h0.dim(1), offset: globalOffset, windowSize: nil) let slidingCache = firstSlidingIdx < cache.count ? cache[firstSlidingIdx] : nil - let slidingMask = createAttentionMask(h: h0, cache: slidingCache, windowSize: config.slidingWindow) + let slidingOffset = slidingCache?.offset ?? 0 + let slidingMask = createAttentionMask(n: h0.dim(1), offset: slidingOffset, windowSize: config.slidingWindow) for i in 0 ..< layers.count { let isGlobal = config.isGlobalLayer(i) diff --git a/packages/swift/Sources/NodeMLXCore/Models/GptOssGenerated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/GptOSSGenerated.swift similarity index 71% rename from packages/swift/Sources/NodeMLXCore/Models/GptOssGenerated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/GptOSSGenerated.swift index a6a9709..c170425 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/GptOssGenerated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/GptOSSGenerated.swift @@ -1,5 +1,5 @@ // -// GptOssGenerated.swift +// GptOSSGenerated.swift // NodeMLXCore // // AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct GptOSSConfiguration: Decodable, Sendable { +public struct GptOSSConfiguration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -99,17 +99,17 @@ public struct GptOSSConfiguration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) - ropeTheta = try decode(.ropeTheta, default: 10000.0) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) + ropeTheta = try decode(.ropeTheta, default: 150_000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: true) mlpBias = try decode(.mlpBias, default: true) // MoE configuration - numLocalExperts = try decode(.numLocalExperts, default: 32) + numLocalExperts = try decode(.numLocalExperts, default: 128) numExpertsPerTok = try decode(.numExpertsPerTok, default: 4) - slidingWindow = try decode(.slidingWindow, default: 512) + slidingWindow = try decode(.slidingWindow, default: 128) slidingWindowPattern = try decode(.slidingWindowPattern, default: 6) if let types: [String] = try? decode(.layerTypes) { @@ -125,24 +125,16 @@ public struct GptOSSConfiguration: Decodable, Sendable { // MARK: - RMS Norm -/// Standard RMSNorm -class GptOSSRMSNorm: Module { - let eps: Float +/// Uses shared RMSNorm implementation +typealias GptOSSRMSNorm = RMSNorm - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } +// MARK: - Utility Functions - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } +/// Top-k selection for MoE routing - uses shared implementation +private func mlxTopK(_ a: MLXArray, k: Int, axis: Int = -1) -> (values: MLXArray, indices: MLXArray) { + MathUtils.topK(a, k: k, axis: axis) } -// MARK: - Utility Functions - // MARK: - Attention class GptOSSAttention: Module { @@ -176,7 +168,7 @@ class GptOSSAttention: Module { _sinks.wrappedValue = MLXArray.zeros([numHeads]) isSliding = !config.isGlobalLayer(layerIdx) let ropeBase = config.ropeTheta - rope = RoPE(dimensions: headDim, traditional: false, base: ropeBase) + rope = RoPE(dimensions: headDim, traditional: true, base: ropeBase) } func callAsFunction( @@ -197,8 +189,8 @@ class GptOSSAttention: Module { // Apply RoPE with cache offset let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) // Update cache if let c = cache { @@ -222,53 +214,39 @@ class GptOSSAttention: Module { // MARK: - MoE MLP -/// Mixture of Experts MLP using shared MoEMLP infrastructure +/// Mixture of Experts MLP with router and experts +/// Uses vendored SwitchLayers from mlx-swift-lm class GptOSSMLP: Module { - @ModuleInfo(key: "router") var router: MoERouter - @ModuleInfo(key: "experts") var experts: SwitchGLU + @ModuleInfo(key: "experts") var experts: SwiGLUSwitchGLU + @ModuleInfo(key: "router") var router: Linear - let numExperts: Int - let topK: Int + let hiddenSize: Int + let numLocalExperts: Int + let numExpertsPerTok: Int init(_ config: GptOSSConfiguration) { - numExperts = config.numLocalExperts - topK = config.numExpertsPerTok + hiddenSize = config.hiddenSize + numLocalExperts = config.numLocalExperts + numExpertsPerTok = config.numExpertsPerTok - _router.wrappedValue = MoERouter( - hiddenSize: config.hiddenSize, - numExperts: config.numLocalExperts, - topK: config.numExpertsPerTok, - bias: config.mlpBias - ) - _experts.wrappedValue = SwitchGLU( + _experts.wrappedValue = SwiGLUSwitchGLU( inputDims: config.hiddenSize, hiddenDims: config.intermediateSize, numExperts: config.numLocalExperts, - bias: config.mlpBias, - useCustomSwiGLU: true + bias: config.mlpBias ) + _router.wrappedValue = Linear(config.hiddenSize, config.numLocalExperts, bias: config.mlpBias) } func callAsFunction(_ x: MLXArray) -> MLXArray { - let shape = x.shape - let batchSeq = shape.dropLast().reduce(1, *) - let hidden = shape.last! - - // Flatten to [batch * seq, hidden] - let xFlat = x.reshaped([batchSeq, hidden]) + let g = router(x) + let (expertScores, indices) = mlxTopK(g, k: numExpertsPerTok, axis: -1) + let expertWeights = softmax(expertScores, axis: -1, precise: true) - // Get routing weights and expert indices - let (weights, indices) = router(xFlat) + var output = experts(x, indices: indices) - // Get expert outputs [batch * seq, topK, hidden] - let expertOutput = experts(xFlat, indices: indices) - - // Weighted sum of expert outputs - let weightsExpanded = weights[.ellipsis, .newAxis] - let weightedOutput = sum(expertOutput * weightsExpanded, axis: 1) - - // Reshape back to original shape - return weightedOutput.reshaped(shape) + output = output * expandedDimensions(expertWeights, axis: -1) + return output.sum(axis: -2) } } @@ -335,9 +313,10 @@ class GptOSSModelInner: Module { if layerType == "full_attention" { firstGlobalIdx = i; break } } let globalCache = firstGlobalIdx < cache.count ? cache[firstGlobalIdx] : nil - let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) - let firstSlidingCache = cache.first ?? nil - let slidingMask = createAttentionMask(h: hiddenStates, cache: firstSlidingCache, windowSize: slidingWindow) + let globalOffset = globalCache?.offset ?? 0 + let globalMask = createAttentionMask(n: hiddenStates.dim(1), offset: globalOffset, windowSize: nil) + let slidingOffset = cache.first??.offset ?? 0 + let slidingMask = createAttentionMask(n: hiddenStates.dim(1), offset: slidingOffset, windowSize: slidingWindow) for i in 0 ..< layers.count { let layerType = i < layerTypes.count ? layerTypes[i] : "sliding_attention" let isGlobal = layerType == "full_attention" @@ -355,70 +334,51 @@ public class GptOSSModel: Module, LLMModel { public let numLayers: Int public let numKVHeads: Int public let headDim: Int + public let kvHeads: [Int] - @ModuleInfo(key: "model") var model: GptOSSModelInner + let model: GptOSSModelInner + private let configuration: GptOSSConfiguration @ModuleInfo(key: "lm_head") var lmHead: Linear - private let config: GptOSSConfiguration - public var supportsCache: Bool { true } public init(_ config: GptOSSConfiguration) { - self.config = config + configuration = config + model = GptOSSModelInner(config) vocabularySize = config.vocabSize numLayers = config.numHiddenLayers numKVHeads = config.numKeyValueHeads headDim = config.headDim - _model.wrappedValue = GptOSSModelInner(config) + kvHeads = (0 ..< config.numHiddenLayers).map { _ in config.numKeyValueHeads } _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) } public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { var cache: [KVCache?] = Array(repeating: nil, count: numLayers) - let h = model(inputIds, cache: &cache) - return lmHead(h) + let hidden = model(inputIds, cache: &cache) + return lmHead(hidden) } public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } else { Array(repeating: nil, count: numLayers) } - let h = model(inputIds, cache: &layerCaches) + let hidden = model(inputIds, cache: &layerCaches) cache = layerCaches.compactMap(\.self) - return lmHead(h) + return lmHead(hidden) } public func newCache() -> [KVCache] { (0 ..< numLayers).map { i in - let layerType = i < config.layerTypes.count ? config.layerTypes[i] : "sliding_attention" + let layerType = i < configuration.layerTypes.count ? configuration.layerTypes[i] : "sliding_attention" if layerType == "full_attention" { return KVCacheSimple() } - else { return RotatingKVCache(maxSize: config.slidingWindow, keep: 0) } + else { return RotatingKVCache(maxSize: configuration.slidingWindow, keep: 0) } } } + // MARK: - Weight Sanitization + + /// Sanitize MoE weights - delegates to shared implementation public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - // Map MoE expert weights to SwitchGLU format - if newKey.contains(".mlp.experts.") { - newKey = newKey.replacingOccurrences(of: ".experts.gate_proj.weight", with: ".experts.gate_proj") - newKey = newKey.replacingOccurrences(of: ".experts.up_proj.weight", with: ".experts.up_proj") - newKey = newKey.replacingOccurrences(of: ".experts.down_proj.weight", with: ".experts.down_proj") - newKey = newKey.replacingOccurrences(of: ".experts.gate_proj.bias", with: ".experts.gate_proj_bias") - newKey = newKey.replacingOccurrences(of: ".experts.up_proj.bias", with: ".experts.up_proj_bias") - newKey = newKey.replacingOccurrences(of: ".experts.down_proj.bias", with: ".experts.down_proj_bias") - } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + MoESanitizer.sanitize(weights: weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift new file mode 100644 index 0000000..cc0754b --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/generated/models/LlamaGenerated.swift @@ -0,0 +1,191 @@ +// +// LlamaGenerated.swift +// NodeMLXCore +// +// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! +// Generated by hf2swift from model patterns. +// Re-run the generator to update this file. +// +// Based on patterns from mlx-lm and mlx-swift-lm. +// + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Configuration + +public struct LlamaConfiguration: Decodable, Sendable, BaseModelConfiguration { + public var hiddenSize: Int + public var numHiddenLayers: Int + public var numAttentionHeads: Int + public var numKeyValueHeads: Int + public var intermediateSize: Int + public var vocabSize: Int + public var headDim: Int + public var rmsNormEps: Float + public var ropeTheta: Float + public var maxPositionEmbeddings: Int + public var attentionBias: Bool + public var mlpBias: Bool + public var ropeScaling: [String: StringOrNumber]? + public var modelType: String? + + enum CodingKeys: String, CodingKey { + case textConfig = "text_config" + case hiddenSize = "hidden_size" + case numHiddenLayers = "num_hidden_layers" + case numAttentionHeads = "num_attention_heads" + case numKeyValueHeads = "num_key_value_heads" + case intermediateSize = "intermediate_size" + case vocabSize = "vocab_size" + case headDim = "head_dim" + case rmsNormEps = "rms_norm_eps" + case ropeTheta = "rope_theta" + case maxPositionEmbeddings = "max_position_embeddings" + case attentionBias = "attention_bias" + case mlpBias = "mlp_bias" + case ropeScaling = "rope_scaling" + case modelType = "model_type" + } + + public init(from decoder: Swift.Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + + // Helper to decode from text_config or top level + func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { + if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), + let value = try? nested.decode(T.self, forKey: key) + { + return value + } + if let value = try? container.decode(T.self, forKey: key) { + return value + } + if let defaultValue { + return defaultValue + } + throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) + } + + hiddenSize = try decode(.hiddenSize) + numHiddenLayers = try decode(.numHiddenLayers) + numAttentionHeads = try decode(.numAttentionHeads) + numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) + + intermediateSize = try decode(.intermediateSize) + + vocabSize = try decode(.vocabSize) + headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) + ropeTheta = try decode(.ropeTheta, default: 10000.0) + maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) + attentionBias = try decode(.attentionBias, default: false) + mlpBias = try decode(.mlpBias, default: false) + + ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) + modelType = try? container.decode(String.self, forKey: .modelType) + } +} + +// MARK: - RMS Norm + +/// Uses shared RMSNorm implementation +typealias LlamaRMSNorm = RMSNorm + +// MARK: - Utility Functions + +// MARK: - Attention + +/// Standard attention - uses shared implementation +typealias LlamaAttention = StandardAttention + +// MARK: - MLP + +/// Standard SwiGLU MLP - uses shared implementation +typealias LlamaMLP = StandardMLP + +// MARK: - Decoder Layer + +/// Standard decoder layer - uses shared implementation +typealias LlamaDecoderLayer = StandardDecoderLayer + +// MARK: - Model Inner + +class LlamaModelInner: Module { + @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding + @ModuleInfo(key: "layers") var layers: [LlamaDecoderLayer] + @ModuleInfo(key: "norm") var norm: LlamaRMSNorm + + let numLayers: Int + let hiddenSize: Int + + init(_ config: LlamaConfiguration) { + numLayers = config.numHiddenLayers + hiddenSize = config.hiddenSize + + _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) + _layers.wrappedValue = (0 ..< numLayers).map { idx in LlamaDecoderLayer(config, layerIdx: idx) } + _norm.wrappedValue = LlamaRMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + } + + func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { + var hiddenStates = embedTokens(inputIds) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) + for i in 0 ..< layers.count { + hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) + } + return norm(hiddenStates) + } +} + +// MARK: - Top-Level Model + +public class LlamaModel: Module, LLMModel { + public let vocabularySize: Int + public let numLayers: Int + public let numKVHeads: Int + public let headDim: Int + + @ModuleInfo(key: "model") var model: LlamaModelInner + @ModuleInfo(key: "lm_head") var lmHead: Linear + + private let config: LlamaConfiguration + + public var supportsCache: Bool { true } + + public init(_ config: LlamaConfiguration) { + self.config = config + vocabularySize = config.vocabSize + numLayers = config.numHiddenLayers + numKVHeads = config.numKeyValueHeads + headDim = config.headDim + _model.wrappedValue = LlamaModelInner(config) + _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) + } + + public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { + var cache: [KVCache?] = Array(repeating: nil, count: numLayers) + let h = model(inputIds, cache: &cache) + return lmHead(h) + } + + public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { + var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } + else { Array(repeating: nil, count: numLayers) } + let h = model(inputIds, cache: &layerCaches) + cache = layerCaches.compactMap(\.self) + return lmHead(h) + } + + public func newCache() -> [KVCache] { + (0 ..< numLayers).map { _ in KVCacheSimple() } + } + + public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { + // Uses shared weight sanitization logic + sanitizeWeights(weights) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/generated/models/Mistral3Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Mistral3Generated.swift new file mode 100644 index 0000000..d1f8263 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Mistral3Generated.swift @@ -0,0 +1,229 @@ +// +// Mistral3Generated.swift +// NodeMLXCore +// +// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! +// Generated by hf2swift from model patterns. +// Re-run the generator to update this file. +// +// Based on patterns from mlx-lm and mlx-swift-lm. +// + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Configuration + +/// YaRN RoPE parameters for long context support +public struct RoPEParameters: Decodable, Sendable { + public var ropeTheta: Float + public var ropeType: String + public var factor: Float + public var mscale: Float + public var mscaleAllDim: Float + public var originalMaxPositionEmbeddings: Int + public var betaFast: Float + public var betaSlow: Float + + enum CodingKeys: String, CodingKey { + case ropeTheta = "rope_theta" + case ropeType = "rope_type" + case factor + case mscale + case mscaleAllDim = "mscale_all_dim" + case originalMaxPositionEmbeddings = "original_max_position_embeddings" + case betaFast = "beta_fast" + case betaSlow = "beta_slow" + } + + public init(from decoder: Swift.Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + ropeTheta = try container.decodeIfPresent(Float.self, forKey: .ropeTheta) ?? 1_000_000.0 + ropeType = try container.decodeIfPresent(String.self, forKey: .ropeType) ?? "yarn" + factor = try container.decodeIfPresent(Float.self, forKey: .factor) ?? 1.0 + mscale = try container.decodeIfPresent(Float.self, forKey: .mscale) ?? 1.0 + mscaleAllDim = try container.decodeIfPresent(Float.self, forKey: .mscaleAllDim) ?? 1.0 + originalMaxPositionEmbeddings = try container.decodeIfPresent(Int.self, forKey: .originalMaxPositionEmbeddings) ?? 16384 + betaFast = try container.decodeIfPresent(Float.self, forKey: .betaFast) ?? 32.0 + betaSlow = try container.decodeIfPresent(Float.self, forKey: .betaSlow) ?? 1.0 + } +} + +public struct Mistral3Configuration: Decodable, Sendable, BaseModelConfiguration { + public var hiddenSize: Int + public var numHiddenLayers: Int + public var numAttentionHeads: Int + public var numKeyValueHeads: Int + public var intermediateSize: Int + public var vocabSize: Int + public var headDim: Int + public var rmsNormEps: Float + public var ropeTheta: Float + public var maxPositionEmbeddings: Int + public var attentionBias: Bool + public var mlpBias: Bool + public var ropeParameters: RoPEParameters? + public var ropeScaling: [String: StringOrNumber]? + public var modelType: String? + + enum CodingKeys: String, CodingKey { + case textConfig = "text_config" + case hiddenSize = "hidden_size" + case numHiddenLayers = "num_hidden_layers" + case numAttentionHeads = "num_attention_heads" + case numKeyValueHeads = "num_key_value_heads" + case intermediateSize = "intermediate_size" + case vocabSize = "vocab_size" + case headDim = "head_dim" + case rmsNormEps = "rms_norm_eps" + case ropeTheta = "rope_theta" + case maxPositionEmbeddings = "max_position_embeddings" + case attentionBias = "attention_bias" + case mlpBias = "mlp_bias" + case ropeParameters = "rope_parameters" + case ropeScaling = "rope_scaling" + case modelType = "model_type" + } + + public init(from decoder: Swift.Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + + // Helper to decode from text_config or top level + func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { + if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), + let value = try? nested.decode(T.self, forKey: key) + { + return value + } + if let value = try? container.decode(T.self, forKey: key) { + return value + } + if let defaultValue { + return defaultValue + } + throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) + } + + hiddenSize = try decode(.hiddenSize) + numHiddenLayers = try decode(.numHiddenLayers) + numAttentionHeads = try decode(.numAttentionHeads) + numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) + + intermediateSize = try decode(.intermediateSize) + + vocabSize = try decode(.vocabSize) + headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) + ropeTheta = try decode(.ropeTheta, default: 1_000_000.0) + maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) + attentionBias = try decode(.attentionBias, default: false) + mlpBias = try decode(.mlpBias, default: false) + + ropeParameters = try? container.decode(RoPEParameters.self, forKey: .ropeParameters) + ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) + modelType = try? container.decode(String.self, forKey: .modelType) + } +} + +// MARK: - RMS Norm + +/// Uses shared RMSNorm implementation +typealias Mistral3RMSNorm = RMSNorm + +// MARK: - Utility Functions + +// MARK: - Attention + +/// Standard attention - uses shared implementation +typealias Mistral3Attention = StandardAttention + +// MARK: - MLP + +/// Standard SwiGLU MLP - uses shared implementation +typealias Mistral3MLP = StandardMLP + +// MARK: - Decoder Layer + +/// Standard decoder layer - uses shared implementation +typealias Mistral3DecoderLayer = StandardDecoderLayer + +// MARK: - Model Inner + +class Mistral3ModelInner: Module { + @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding + @ModuleInfo(key: "layers") var layers: [Mistral3DecoderLayer] + @ModuleInfo(key: "norm") var norm: Mistral3RMSNorm + + let numLayers: Int + let hiddenSize: Int + + init(_ config: Mistral3Configuration) { + numLayers = config.numHiddenLayers + hiddenSize = config.hiddenSize + + _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) + _layers.wrappedValue = (0 ..< numLayers).map { idx in Mistral3DecoderLayer(config, layerIdx: idx) } + _norm.wrappedValue = Mistral3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + } + + func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { + var hiddenStates = embedTokens(inputIds) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) + for i in 0 ..< layers.count { + hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) + } + return norm(hiddenStates) + } +} + +// MARK: - Top-Level Model + +public class Mistral3Model: Module, LLMModel { + public let vocabularySize: Int + public let numLayers: Int + public let numKVHeads: Int + public let headDim: Int + + @ModuleInfo(key: "model") var model: Mistral3ModelInner + @ModuleInfo(key: "lm_head") var lmHead: Linear + + private let config: Mistral3Configuration + + public var supportsCache: Bool { true } + + public init(_ config: Mistral3Configuration) { + self.config = config + vocabularySize = config.vocabSize + numLayers = config.numHiddenLayers + numKVHeads = config.numKeyValueHeads + headDim = config.headDim + _model.wrappedValue = Mistral3ModelInner(config) + _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) + } + + public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { + var cache: [KVCache?] = Array(repeating: nil, count: numLayers) + let h = model(inputIds, cache: &cache) + return lmHead(h) + } + + public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { + var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } + else { Array(repeating: nil, count: numLayers) } + let h = model(inputIds, cache: &layerCaches) + cache = layerCaches.compactMap(\.self) + return lmHead(h) + } + + public func newCache() -> [KVCache] { + (0 ..< numLayers).map { _ in KVCacheSimple() } + } + + public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { + // Uses shared weight sanitization logic + sanitizeWeights(weights) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/Models/MistralGenerated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/MistralGenerated.swift similarity index 81% rename from packages/swift/Sources/NodeMLXCore/Models/MistralGenerated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/MistralGenerated.swift index 22c6b11..844e8d7 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/MistralGenerated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/MistralGenerated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct MistralConfiguration: Decodable, Sendable { +public struct MistralConfiguration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -87,7 +87,7 @@ public struct MistralConfiguration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) ropeTheta = try decode(.ropeTheta, default: 10000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -102,21 +102,8 @@ public struct MistralConfiguration: Decodable, Sendable { // MARK: - RMS Norm -/// Standard RMSNorm -class MistralRMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} +/// Uses shared RMSNorm implementation +typealias MistralRMSNorm = RMSNorm // MARK: - Utility Functions @@ -172,8 +159,8 @@ class MistralAttention: Module { // Apply RoPE with cache offset let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) // Update cache if let c = cache { @@ -197,23 +184,8 @@ class MistralAttention: Module { // MARK: - MLP -class MistralMLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: MistralConfiguration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} +/// Standard SwiGLU MLP - uses shared implementation +typealias MistralMLP = StandardMLP // MARK: - Decoder Layer @@ -274,11 +246,12 @@ class MistralModelInner: Module { var hiddenStates = embedTokens(inputIds) let globalLayerIdx = slidingWindowPattern - 1 let globalCache = globalLayerIdx < cache.count ? cache[globalLayerIdx] : nil - let globalMask = createAttentionMask(h: hiddenStates, cache: globalCache, windowSize: nil) + let globalOffset = globalCache?.offset ?? 0 + let globalMask = createAttentionMask(n: hiddenStates.dim(1), offset: globalOffset, windowSize: nil) let slidingMask: MLXFast.ScaledDotProductAttentionMaskMode if slidingWindowPattern > 1 { - let firstCache = cache.first ?? nil - slidingMask = createAttentionMask(h: hiddenStates, cache: firstCache, windowSize: slidingWindow) + let slidingOffset = cache.first??.offset ?? 0 + slidingMask = createAttentionMask(n: hiddenStates.dim(1), offset: slidingOffset, windowSize: slidingWindow) } else { slidingMask = globalMask } @@ -339,20 +312,7 @@ public class MistralModel: Module, LLMModel { } public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + // Uses shared weight sanitization logic + sanitizeWeights(weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/Models/Phi3Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Phi3Generated.swift similarity index 56% rename from packages/swift/Sources/NodeMLXCore/Models/Phi3Generated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/Phi3Generated.swift index 1ea9bb0..3b9ad9f 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/Phi3Generated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Phi3Generated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct Phi3Configuration: Decodable, Sendable { +public struct Phi3Configuration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -78,7 +78,7 @@ public struct Phi3Configuration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) ropeTheta = try decode(.ropeTheta, default: 10000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -91,95 +91,18 @@ public struct Phi3Configuration: Decodable, Sendable { // MARK: - RMS Norm -/// Standard RMSNorm -class Phi3RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} +/// Uses shared RMSNorm implementation +typealias Phi3RMSNorm = RMSNorm // MARK: - Utility Functions // MARK: - Attention -class Phi3Attention: Module { - @ModuleInfo(key: "qkv_proj") var qkvProj: Linear - @ModuleInfo(key: "o_proj") var oProj: Linear - - let numHeads: Int - let numKVHeads: Int - let headDim: Int - let scale: Float - let rope: RoPE - - init(_ config: Phi3Configuration) { - numHeads = config.numAttentionHeads - numKVHeads = config.numKeyValueHeads - headDim = config.headDim - scale = 1.0 / sqrt(Float(headDim)) - - let qDim = numHeads * headDim - let kvDim = numKVHeads * headDim - let opSize = qDim + 2 * kvDim +/// AttentionConfiguration conformance for fused QKV attention +extension Phi3Configuration: AttentionConfiguration {} - _qkvProj.wrappedValue = Linear(config.hiddenSize, opSize, bias: false) - _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: false) - rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) - - let qkv = qkvProj(hiddenStates) - let queryPos = numHeads * headDim - let kvPos = queryPos + numKVHeads * headDim - - var queries = qkv[0..., 0..., .. [B, L, hidden] - let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) - return oProj(outputReshaped) - } -} +/// Fused QKV attention - uses shared implementation +typealias Phi3Attention = FusedQKVAttention // MARK: - MLP @@ -204,36 +127,8 @@ class Phi3MLP: Module { // MARK: - Decoder Layer -class Phi3DecoderLayer: Module { - @ModuleInfo(key: "self_attn") var selfAttn: Phi3Attention - @ModuleInfo(key: "mlp") var mlp: Phi3MLP - @ModuleInfo(key: "input_layernorm") var inputLayernorm: Phi3RMSNorm - @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: Phi3RMSNorm - - init(_ config: Phi3Configuration, layerIdx _: Int = 0) { - _selfAttn.wrappedValue = Phi3Attention(config) - _mlp.wrappedValue = Phi3MLP(config) - _inputLayernorm.wrappedValue = Phi3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - _postAttentionLayernorm.wrappedValue = Phi3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - } - - func callAsFunction( - _ hiddenStates: MLXArray, - mask: MLXFast.ScaledDotProductAttentionMaskMode, - cache: inout KVCache? - ) -> MLXArray { - // 1. Pre-norm + Self-attention - let normed = inputLayernorm(hiddenStates) - let attnOut = selfAttn(normed, mask: mask, cache: &cache) - var h = hiddenStates + attnOut - - // 2. Pre-norm + MLP - let mlpNormed = postAttentionLayernorm(h) - let mlpOut = mlp(mlpNormed) - h = h + mlpOut - return h - } -} +/// Standard decoder layer - uses shared implementation +typealias Phi3DecoderLayer = StandardDecoderLayer // MARK: - Model Inner @@ -256,7 +151,8 @@ class Phi3ModelInner: Module { func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { var hiddenStates = embedTokens(inputIds) - let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) for i in 0 ..< layers.count { hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) } @@ -308,20 +204,7 @@ public class Phi3Model: Module, LLMModel { } public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + // Uses shared weight sanitization logic + sanitizeWeights(weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/generated/models/Qwen2Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Qwen2Generated.swift new file mode 100644 index 0000000..c23c17f --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Qwen2Generated.swift @@ -0,0 +1,191 @@ +// +// Qwen2Generated.swift +// NodeMLXCore +// +// AUTO-GENERATED FILE - DO NOT EDIT MANUALLY! +// Generated by hf2swift from model patterns. +// Re-run the generator to update this file. +// +// Based on patterns from mlx-lm and mlx-swift-lm. +// + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Configuration + +public struct Qwen2Configuration: Decodable, Sendable, BaseModelConfiguration { + public var hiddenSize: Int + public var numHiddenLayers: Int + public var numAttentionHeads: Int + public var numKeyValueHeads: Int + public var intermediateSize: Int + public var vocabSize: Int + public var headDim: Int + public var rmsNormEps: Float + public var ropeTheta: Float + public var maxPositionEmbeddings: Int + public var attentionBias: Bool + public var mlpBias: Bool + public var ropeScaling: [String: StringOrNumber]? + public var modelType: String? + + enum CodingKeys: String, CodingKey { + case textConfig = "text_config" + case hiddenSize = "hidden_size" + case numHiddenLayers = "num_hidden_layers" + case numAttentionHeads = "num_attention_heads" + case numKeyValueHeads = "num_key_value_heads" + case intermediateSize = "intermediate_size" + case vocabSize = "vocab_size" + case headDim = "head_dim" + case rmsNormEps = "rms_norm_eps" + case ropeTheta = "rope_theta" + case maxPositionEmbeddings = "max_position_embeddings" + case attentionBias = "attention_bias" + case mlpBias = "mlp_bias" + case ropeScaling = "rope_scaling" + case modelType = "model_type" + } + + public init(from decoder: Swift.Decoder) throws { + let container = try decoder.container(keyedBy: CodingKeys.self) + + // Helper to decode from text_config or top level + func decode(_ key: CodingKeys, default defaultValue: T? = nil) throws -> T { + if let nested = try? container.nestedContainer(keyedBy: CodingKeys.self, forKey: .textConfig), + let value = try? nested.decode(T.self, forKey: key) + { + return value + } + if let value = try? container.decode(T.self, forKey: key) { + return value + } + if let defaultValue { + return defaultValue + } + throw DecodingError.keyNotFound(key, DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)")) + } + + hiddenSize = try decode(.hiddenSize) + numHiddenLayers = try decode(.numHiddenLayers) + numAttentionHeads = try decode(.numAttentionHeads) + numKeyValueHeads = try decode(.numKeyValueHeads, default: numAttentionHeads) + + intermediateSize = try decode(.intermediateSize) + + vocabSize = try decode(.vocabSize) + headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) + rmsNormEps = try decode(.rmsNormEps, default: 0.000001) + ropeTheta = try decode(.ropeTheta, default: 10000.0) + maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) + attentionBias = try decode(.attentionBias, default: true) + mlpBias = try decode(.mlpBias, default: false) + + ropeScaling = try? container.decode([String: StringOrNumber].self, forKey: .ropeScaling) + modelType = try? container.decode(String.self, forKey: .modelType) + } +} + +// MARK: - RMS Norm + +/// Uses shared RMSNorm implementation +typealias Qwen2RMSNorm = RMSNorm + +// MARK: - Utility Functions + +// MARK: - Attention + +/// Standard attention - uses shared implementation +typealias Qwen2Attention = StandardAttention + +// MARK: - MLP + +/// Standard SwiGLU MLP - uses shared implementation +typealias Qwen2MLP = StandardMLP + +// MARK: - Decoder Layer + +/// Standard decoder layer - uses shared implementation +typealias Qwen2DecoderLayer = StandardDecoderLayer + +// MARK: - Model Inner + +class Qwen2ModelInner: Module { + @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding + @ModuleInfo(key: "layers") var layers: [Qwen2DecoderLayer] + @ModuleInfo(key: "norm") var norm: Qwen2RMSNorm + + let numLayers: Int + let hiddenSize: Int + + init(_ config: Qwen2Configuration) { + numLayers = config.numHiddenLayers + hiddenSize = config.hiddenSize + + _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) + _layers.wrappedValue = (0 ..< numLayers).map { idx in Qwen2DecoderLayer(config, layerIdx: idx) } + _norm.wrappedValue = Qwen2RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + } + + func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { + var hiddenStates = embedTokens(inputIds) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) + for i in 0 ..< layers.count { + hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) + } + return norm(hiddenStates) + } +} + +// MARK: - Top-Level Model + +public class Qwen2Model: Module, LLMModel { + public let vocabularySize: Int + public let numLayers: Int + public let numKVHeads: Int + public let headDim: Int + + @ModuleInfo(key: "model") var model: Qwen2ModelInner + @ModuleInfo(key: "lm_head") var lmHead: Linear + + private let config: Qwen2Configuration + + public var supportsCache: Bool { true } + + public init(_ config: Qwen2Configuration) { + self.config = config + vocabularySize = config.vocabSize + numLayers = config.numHiddenLayers + numKVHeads = config.numKeyValueHeads + headDim = config.headDim + _model.wrappedValue = Qwen2ModelInner(config) + _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) + } + + public func callAsFunction(_ inputIds: MLXArray) -> MLXArray { + var cache: [KVCache?] = Array(repeating: nil, count: numLayers) + let h = model(inputIds, cache: &cache) + return lmHead(h) + } + + public func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache]?) -> MLXArray { + var layerCaches: [KVCache?] = if let existingCache = cache { existingCache.map { $0 as KVCache? } } + else { Array(repeating: nil, count: numLayers) } + let h = model(inputIds, cache: &layerCaches) + cache = layerCaches.compactMap(\.self) + return lmHead(h) + } + + public func newCache() -> [KVCache] { + (0 ..< numLayers).map { _ in KVCacheSimple() } + } + + public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { + // Uses shared weight sanitization logic + sanitizeWeights(weights) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/Models/Qwen3Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/Qwen3Generated.swift similarity index 80% rename from packages/swift/Sources/NodeMLXCore/Models/Qwen3Generated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/Qwen3Generated.swift index e4e029e..d43ff8a 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/Qwen3Generated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/Qwen3Generated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct Qwen3Configuration: Decodable, Sendable { +public struct Qwen3Configuration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -78,7 +78,7 @@ public struct Qwen3Configuration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.000001) ropeTheta = try decode(.ropeTheta, default: 1_000_000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -91,21 +91,8 @@ public struct Qwen3Configuration: Decodable, Sendable { // MARK: - RMS Norm -/// Standard RMSNorm -class Qwen3RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} +/// Uses shared RMSNorm implementation +typealias Qwen3RMSNorm = RMSNorm // MARK: - Utility Functions @@ -164,8 +151,8 @@ class Qwen3Attention: Module { // Apply RoPE with cache offset let offset = cache?.offset ?? 0 - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) // Update cache if let c = cache { @@ -189,23 +176,8 @@ class Qwen3Attention: Module { // MARK: - MLP -class Qwen3MLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: Qwen3Configuration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} +/// Standard SwiGLU MLP - uses shared implementation +typealias Qwen3MLP = StandardMLP // MARK: - Decoder Layer @@ -261,7 +233,8 @@ class Qwen3ModelInner: Module { func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { var hiddenStates = embedTokens(inputIds) - let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) for i in 0 ..< layers.count { hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) } @@ -313,20 +286,7 @@ public class Qwen3Model: Module, LLMModel { } public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + // Uses shared weight sanitization logic + sanitizeWeights(weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/Models/SmolLM3Generated.swift b/packages/swift/Sources/NodeMLXCore/generated/models/SmolLM3Generated.swift similarity index 72% rename from packages/swift/Sources/NodeMLXCore/Models/SmolLM3Generated.swift rename to packages/swift/Sources/NodeMLXCore/generated/models/SmolLM3Generated.swift index 218c899..766edf5 100644 --- a/packages/swift/Sources/NodeMLXCore/Models/SmolLM3Generated.swift +++ b/packages/swift/Sources/NodeMLXCore/generated/models/SmolLM3Generated.swift @@ -16,7 +16,7 @@ import MLXNN // MARK: - Configuration -public struct Smollm3Configuration: Decodable, Sendable { +public struct SmolLM3Configuration: Decodable, Sendable, BaseModelConfiguration { public var hiddenSize: Int public var numHiddenLayers: Int public var numAttentionHeads: Int @@ -88,7 +88,7 @@ public struct Smollm3Configuration: Decodable, Sendable { vocabSize = try decode(.vocabSize) headDim = try decode(.headDim, default: hiddenSize / numAttentionHeads) - rmsNormEps = try decode(.rmsNormEps, default: 1e-6) + rmsNormEps = try decode(.rmsNormEps, default: 0.00001) ropeTheta = try decode(.ropeTheta, default: 5_000_000.0) maxPositionEmbeddings = try decode(.maxPositionEmbeddings, default: 32768) attentionBias = try decode(.attentionBias, default: false) @@ -109,27 +109,14 @@ public struct Smollm3Configuration: Decodable, Sendable { // MARK: - RMS Norm -/// Standard RMSNorm -class Smollm3RMSNorm: Module { - let eps: Float - - @ModuleInfo(key: "weight") var weight: MLXArray - - init(dimensions: Int, eps: Float = 1e-6) { - self.eps = eps - _weight.wrappedValue = MLXArray.ones([dimensions]) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - MLXFast.rmsNorm(x, weight: weight, eps: eps) - } -} +/// Uses shared RMSNorm implementation +typealias SmolLM3RMSNorm = RMSNorm // MARK: - Utility Functions // MARK: - Attention -class Smollm3Attention: Module { +class SmolLM3Attention: Module { @ModuleInfo(key: "q_proj") var qProj: Linear @ModuleInfo(key: "k_proj") var kProj: Linear @ModuleInfo(key: "v_proj") var vProj: Linear @@ -142,7 +129,7 @@ class Smollm3Attention: Module { let rope: RoPE let skipRope: Bool - init(_ config: Smollm3Configuration, layerIdx: Int) { + init(_ config: SmolLM3Configuration, layerIdx: Int) { numHeads = config.numAttentionHeads numKVHeads = config.numKeyValueHeads headDim = config.headDim @@ -179,8 +166,8 @@ class Smollm3Attention: Module { // Apply RoPE with cache offset let offset = cache?.offset ?? 0 if !skipRope { - queries = rope.apply(queries, offset: offset) - keys = rope.apply(keys, offset: offset) + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) } // Update cache @@ -205,37 +192,22 @@ class Smollm3Attention: Module { // MARK: - MLP -class Smollm3MLP: Module { - @ModuleInfo(key: "gate_proj") var gateProj: Linear - @ModuleInfo(key: "up_proj") var upProj: Linear - @ModuleInfo(key: "down_proj") var downProj: Linear - - init(_ config: Smollm3Configuration) { - let intermediateSize = config.intermediateSize - let mlpBias = config.mlpBias - _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) - _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) - } - - func callAsFunction(_ x: MLXArray) -> MLXArray { - downProj(silu(gateProj(x)) * upProj(x)) - } -} +/// Standard SwiGLU MLP - uses shared implementation +typealias SmolLM3MLP = StandardMLP // MARK: - Decoder Layer -class Smollm3DecoderLayer: Module { - @ModuleInfo(key: "self_attn") var selfAttn: Smollm3Attention - @ModuleInfo(key: "mlp") var mlp: Smollm3MLP - @ModuleInfo(key: "input_layernorm") var inputLayernorm: Smollm3RMSNorm - @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: Smollm3RMSNorm - - init(_ config: Smollm3Configuration, layerIdx: Int) { - _selfAttn.wrappedValue = Smollm3Attention(config, layerIdx: layerIdx) - _mlp.wrappedValue = Smollm3MLP(config) - _inputLayernorm.wrappedValue = Smollm3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) - _postAttentionLayernorm.wrappedValue = Smollm3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) +class SmolLM3DecoderLayer: Module { + @ModuleInfo(key: "self_attn") var selfAttn: SmolLM3Attention + @ModuleInfo(key: "mlp") var mlp: SmolLM3MLP + @ModuleInfo(key: "input_layernorm") var inputLayernorm: SmolLM3RMSNorm + @ModuleInfo(key: "post_attention_layernorm") var postAttentionLayernorm: SmolLM3RMSNorm + + init(_ config: SmolLM3Configuration, layerIdx: Int) { + _selfAttn.wrappedValue = SmolLM3Attention(config, layerIdx: layerIdx) + _mlp.wrappedValue = SmolLM3MLP(config) + _inputLayernorm.wrappedValue = SmolLM3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + _postAttentionLayernorm.wrappedValue = SmolLM3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) } func callAsFunction( @@ -258,26 +230,27 @@ class Smollm3DecoderLayer: Module { // MARK: - Model Inner -class Smollm3ModelInner: Module { +class SmolLM3ModelInner: Module { @ModuleInfo(key: "embed_tokens") var embedTokens: Embedding - @ModuleInfo(key: "layers") var layers: [Smollm3DecoderLayer] - @ModuleInfo(key: "norm") var norm: Smollm3RMSNorm + @ModuleInfo(key: "layers") var layers: [SmolLM3DecoderLayer] + @ModuleInfo(key: "norm") var norm: SmolLM3RMSNorm let numLayers: Int let hiddenSize: Int - init(_ config: Smollm3Configuration) { + init(_ config: SmolLM3Configuration) { numLayers = config.numHiddenLayers hiddenSize = config.hiddenSize _embedTokens.wrappedValue = Embedding(embeddingCount: config.vocabSize, dimensions: config.hiddenSize) - _layers.wrappedValue = (0 ..< numLayers).map { idx in Smollm3DecoderLayer(config, layerIdx: idx) } - _norm.wrappedValue = Smollm3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + _layers.wrappedValue = (0 ..< numLayers).map { idx in SmolLM3DecoderLayer(config, layerIdx: idx) } + _norm.wrappedValue = SmolLM3RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) } func callAsFunction(_ inputIds: MLXArray, cache: inout [KVCache?]) -> MLXArray { var hiddenStates = embedTokens(inputIds) - let mask = createAttentionMask(h: hiddenStates, cache: cache.first ?? nil, windowSize: nil) + let offset = cache.first??.offset ?? 0 + let mask = createAttentionMask(n: hiddenStates.dim(1), offset: offset, windowSize: nil) for i in 0 ..< layers.count { hiddenStates = layers[i](hiddenStates, mask: mask, cache: &cache[i]) } @@ -287,26 +260,26 @@ class Smollm3ModelInner: Module { // MARK: - Top-Level Model -public class Smollm3Model: Module, LLMModel { +public class SmolLM3Model: Module, LLMModel { public let vocabularySize: Int public let numLayers: Int public let numKVHeads: Int public let headDim: Int - @ModuleInfo(key: "model") var model: Smollm3ModelInner + @ModuleInfo(key: "model") var model: SmolLM3ModelInner @ModuleInfo(key: "lm_head") var lmHead: Linear - private let config: Smollm3Configuration + private let config: SmolLM3Configuration public var supportsCache: Bool { true } - public init(_ config: Smollm3Configuration) { + public init(_ config: SmolLM3Configuration) { self.config = config vocabularySize = config.vocabSize numLayers = config.numHiddenLayers numKVHeads = config.numKeyValueHeads headDim = config.headDim - _model.wrappedValue = Smollm3ModelInner(config) + _model.wrappedValue = SmolLM3ModelInner(config) _lmHead.wrappedValue = Linear(config.hiddenSize, config.vocabSize, bias: false) } @@ -329,20 +302,7 @@ public class Smollm3Model: Module, LLMModel { } public func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { - var result: [String: MLXArray] = [:] - for (key, value) in weights { - var newKey = key - if newKey.hasPrefix("language_model.model.") { newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) } - else if newKey.hasPrefix("language_model.lm_head.") { newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) } - else if newKey.hasPrefix("language_model.") { newKey = String(newKey.dropFirst("language_model.".count)) } - if newKey.contains("vision_tower") || newKey.contains("audio_tower") || newKey.contains("multi_modal_projector") { continue } - result[newKey] = value - } - if result["lm_head.weight"] == nil { - for suffix in ["weight", "scales", "biases"] { - if let embedWeight = result["model.embed_tokens.\(suffix)"] { result["lm_head.\(suffix)"] = embedWeight } - } - } - return result + // Uses shared weight sanitization logic + sanitizeWeights(weights) } } diff --git a/packages/swift/Sources/NodeMLXCore/ported/GemmaRMSNorm.swift b/packages/swift/Sources/NodeMLXCore/ported/GemmaRMSNorm.swift new file mode 100644 index 0000000..a592e03 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/ported/GemmaRMSNorm.swift @@ -0,0 +1,46 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/gemma.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import MLX +import MLXFast +import MLXNN + +/// Gemma-style RMSNorm with (1 + weight) scaling. +/// +/// Unlike standard RMSNorm which uses `weight` directly, Gemma models +/// use `(1 + weight)` scaling. This means the weight is initialized to +/// zeros and the effective scale is `1 + weight`. +/// +/// This is used by Gemma, Gemma2, Gemma3, and Gemma3n models. +/// +/// Original Python: +/// ```python +/// class RMSNorm(nn.Module): +/// def __init__(self, dims: int, eps: float = 1e-5): +/// super().__init__() +/// self.weight = mx.ones((dims,)) +/// self.eps = eps +/// +/// def __call__(self, x): +/// return mx.fast.rms_norm(x, 1.0 + self.weight, self.eps) +/// ``` +public class GemmaRMSNorm: Module { + public let eps: Float + + @ModuleInfo(key: "weight") public var weight: MLXArray + + public init(dimensions: Int, eps: Float = 1e-5) { + self.eps = eps + // Initialize to zeros - effective scale will be (1 + weight) = 1 + _weight.wrappedValue = MLXArray.zeros([dimensions]) + } + + public func callAsFunction(_ x: MLXArray) -> MLXArray { + // Gemma uses (1 + weight) scaling + MLXFast.rmsNorm(x, weight: 1 + weight, eps: eps) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/ported/KVCache.swift b/packages/swift/Sources/NodeMLXCore/ported/KVCache.swift new file mode 100644 index 0000000..8f09930 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/ported/KVCache.swift @@ -0,0 +1,643 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/cache.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Causal Mask Creation + +/// Creates a causal attention mask for autoregressive decoding. +/// +/// The mask ensures that each position can only attend to itself and previous positions. +/// +/// - Parameters: +/// - n: Query sequence length +/// - offset: Number of previously cached tokens +/// - windowSize: Optional sliding window size for local attention +/// - Returns: Causal mask as MLXArray with shape [n, offset + n] +public func createCausalMask( + n: Int, + offset: Int = 0, + windowSize: Int? = nil +) -> MLXArray { + // Row indices: [0, 1, ..., n-1] + offset + let rowIndices = MLXArray(Int32(offset) ..< Int32(offset + n)) + .reshaped([n, 1]) + + // Column indices: [0, 1, ..., offset + n - 1] + let colIndices = MLXArray(0 ..< Int32(offset + n)) + .reshaped([1, offset + n]) + + // Causal: can only attend to current and previous positions + var mask = rowIndices .>= colIndices + + // Optional window constraint: can only attend within window + if let windowSize { + let windowMask = rowIndices .< (colIndices + Int32(windowSize)) + mask = logicalAnd(mask, windowMask) + } + + return mask +} + +/// Creates an attention mask appropriate for the given parameters. +/// +/// - Parameters: +/// - n: Query sequence length +/// - offset: Cache offset (number of previously cached tokens) +/// - returnArray: If true, always returns array mask; if false, may return "causal" string +/// - windowSize: Optional sliding window size +/// - Returns: Mask mode for MLXFast scaled dot product attention +public func createAttentionMask( + n: Int, + offset: Int, + returnArray: Bool = false, + windowSize: Int? = nil +) -> MLXFast.ScaledDotProductAttentionMaskMode { + // Single token generation with no window constraint - no mask needed + if n == 1 && windowSize == nil { + return .none + } + + if returnArray || windowSize != nil { + return .array(createCausalMask(n: n, offset: offset, windowSize: windowSize)) + } else { + return .causal + } +} + +// MARK: - KVCache Protocol + +/// Protocol for all KV cache implementations. +/// +/// Caches store key-value pairs from previous forward passes to enable +/// efficient autoregressive generation without recomputing attention +/// over the entire sequence. +public protocol KVCacheProtocol: AnyObject { + /// Updates the cache with new keys/values and returns the full sequence. + /// + /// - Parameters: + /// - keys: New keys to add, shape [B, H, S, D] + /// - values: New values to add, shape [B, H, S, D] + /// - Returns: Tuple of (allKeys, allValues) including new and cached entries + func update(keys: MLXArray, values: MLXArray) -> (MLXArray, MLXArray) + + /// Number of cached tokens. + var offset: Int { get } + + /// Current cached keys and values (for KV-sharing scenarios like Gemma3n). + /// Returns nil if no cache exists yet. + var state: (keys: MLXArray, values: MLXArray)? { get } + + /// Whether this cache can be trimmed. + var isTrimmable: Bool { get } + + /// Trim the cache by removing the last n tokens. + /// - Returns: Actual number of tokens trimmed + @discardableResult + func trim(_ n: Int) -> Int + + /// Create attention mask for the current cache state. + /// + /// - Parameters: + /// - queryLength: Length of the query sequence + /// - windowSize: Optional sliding window size + /// - returnArray: If true, always returns array mask + /// - Returns: Mask mode for scaled dot product attention + func makeMask( + queryLength: Int, + windowSize: Int?, + returnArray: Bool + ) -> MLXFast.ScaledDotProductAttentionMaskMode +} + +// MARK: - Default Protocol Implementation + +public extension KVCacheProtocol { + var isTrimmable: Bool { true } + + func makeMask( + queryLength: Int, + windowSize: Int? = nil, + returnArray: Bool = false + ) -> MLXFast.ScaledDotProductAttentionMaskMode { + createAttentionMask( + n: queryLength, + offset: offset, + returnArray: returnArray, + windowSize: windowSize + ) + } +} + +// MARK: - KVCache + +/// Standard KV cache with grow-in-place strategy for efficient memory use. +/// +/// Uses a step-based allocation strategy to avoid frequent reallocations. +/// The internal buffer grows in steps of `step` (256) tokens. +/// +/// Ported from: mlx_lm/models/cache.py::KVCache +public final class StandardKVCache: KVCacheProtocol { + /// Growth step size for buffer allocation + public static let step = 256 + + private var keys: MLXArray? + private var values: MLXArray? + public private(set) var offset: Int = 0 + + public init() {} + + /// Returns the current cached keys and values. + public var state: (keys: MLXArray, values: MLXArray)? { + guard let k = keys, let v = values, offset > 0 else { return nil } + return (k[.ellipsis, .. (MLXArray, MLXArray) { + let prev = offset + let numSteps = newKeys.dim(2) + + // Check if we need to grow the buffer + if keys == nil || (prev + numSteps) > keys!.dim(2) { + let batchSize = newKeys.dim(0) + let numKvHeads = newKeys.dim(1) + let keyHeadDim = newKeys.dim(3) + let valueHeadDim = newValues.dim(3) + + // Calculate new buffer size (round up to step boundary) + let nBufferSteps = (Self.step + numSteps - 1) / Self.step + let bufferSize = nBufferSteps * Self.step + + let kShape = [batchSize, numKvHeads, bufferSize, keyHeadDim] + let vShape = [batchSize, numKvHeads, bufferSize, valueHeadDim] + + let newK = MLXArray.zeros(kShape, dtype: newKeys.dtype) + let newV = MLXArray.zeros(vShape, dtype: newValues.dtype) + + if let existingKeys = keys, let existingValues = values { + // Trim existing buffer if not aligned to step + var trimmedKeys = existingKeys + var trimmedValues = existingValues + if prev % Self.step != 0 { + trimmedKeys = existingKeys[.ellipsis, .. Int { + let trimmed = min(offset, n) + offset -= trimmed + return trimmed + } + + /// Converts this cache to a quantized version. + /// + /// - Parameters: + /// - groupSize: Quantization group size (default: 64) + /// - bits: Bits per weight (default: 4) + /// - Returns: New QuantizedKVCache with quantized contents + public func toQuantized(groupSize: Int = 64, bits: Int = 4) -> QuantizedKVCache { + let quantCache = QuantizedKVCache(groupSize: groupSize, bits: bits) + quantCache.offset = offset + if let k = keys, let v = values { + quantCache.keys = MLX.quantized(k, groupSize: groupSize, bits: bits) + quantCache.values = MLX.quantized(v, groupSize: groupSize, bits: bits) + } + return quantCache + } +} + +// MARK: - Compatibility Alias + +/// Alias for backward compatibility - use StandardKVCache directly for new code. +@available(*, deprecated, renamed: "StandardKVCache") +public typealias SimpleKVCache = StandardKVCache + +// MARK: - RotatingKVCache + +/// Rotating KV cache for sliding window attention. +/// +/// Maintains a fixed-size window of the most recent tokens, with optional +/// "attention sinks" (kept tokens at the beginning) for stability. +/// +/// Ported from: mlx_lm/models/cache.py::RotatingKVCache +public final class RotatingKVCache: KVCacheProtocol { + /// Growth step size for buffer allocation + public static let step = 256 + + /// Number of initial tokens to keep as attention sinks + public let keep: Int + + /// Maximum cache size (sliding window size) + public let maxSize: Int + + private var keys: MLXArray? + private var values: MLXArray? + public private(set) var offset: Int = 0 + + /// Internal write index for rotation + private var idx: Int = 0 + + /// Creates a rotating KV cache. + /// + /// - Parameters: + /// - maxSize: Maximum number of tokens to keep in cache + /// - keep: Number of initial tokens to preserve as attention sinks (default: 0) + public init(maxSize: Int, keep: Int = 0) { + self.maxSize = maxSize + self.keep = keep + } + + /// Returns the current cached keys and values in temporal order. + public var state: (keys: MLXArray, values: MLXArray)? { + guard let k = keys, let v = values else { return nil } + let reorderedK = temporalOrder(k) + let reorderedV = temporalOrder(v) + return (reorderedK, reorderedV) + } + + // MARK: - Private Helpers + + /// Trims the cache and optionally appends new values. + private func trimBuffer(_ trimSize: Int, _ v: MLXArray, append: MLXArray? = nil) -> MLXArray { + var toCat: [MLXArray] = [] + + if trimSize > 0 { + // Keep the "sink" tokens and skip trimmed portion + toCat.append(v[.ellipsis, .. MLXArray { + if idx == v.dim(2) { + v + } else if idx < offset { + // Cache has rotated - reorder + concatenated([ + v[.ellipsis, .. (MLXArray, MLXArray) { + if keys == nil { + keys = newKeys + values = newValues + } else { + // Reorder to temporal order to preserve context + keys = temporalOrder(keys!) + values = temporalOrder(values!) + idx = keys!.dim(2) + + // Calculate trim size (keep at least maxSize context) + let trimSize = idx - maxSize + 1 + keys = trimBuffer(trimSize, keys!, append: newKeys) + values = trimBuffer(trimSize, values!, append: newValues) + } + + offset += newKeys.dim(2) + idx = keys!.dim(2) + return (keys!, values!) + } + + /// Update in-place (for single-token generation). + private func updateInPlace(keys newKeys: MLXArray, values newValues: MLXArray) -> (MLXArray, MLXArray) { + let batchSize = newKeys.dim(0) + let numKvHeads = newKeys.dim(1) + let seqLen = newKeys.dim(2) + let keyHeadDim = newKeys.dim(3) + let valueHeadDim = newValues.dim(3) + + let prev = offset + + // Grow cache if needed (up to maxSize) + if keys == nil || (prev >= keys!.dim(2) && keys!.dim(2) < maxSize) { + let newSize = min(Self.step, maxSize - prev) + let kShape = [batchSize, numKvHeads, newSize, keyHeadDim] + let vShape = [batchSize, numKvHeads, newSize, valueHeadDim] + + let newK = MLXArray.zeros(kShape, dtype: newKeys.dtype) + let newV = MLXArray.zeros(vShape, dtype: newValues.dtype) + + if let existingKeys = keys, let existingValues = values { + keys = concatenated([existingKeys, newK], axis: 2) + values = concatenated([existingValues, newV], axis: 2) + } else { + keys = newK + values = newV + } + idx = prev + } + + // Trim if needed + let trimSize = keys!.dim(2) - maxSize + if trimSize > 0 { + keys = trimBuffer(trimSize, keys!) + values = trimBuffer(trimSize, values!) + idx = maxSize + } + + // Rotate index when we hit max size + if idx == maxSize { + idx = keep + } + + // Assign new values at current position + keys![.ellipsis, idx ..< (idx + seqLen), 0...] = newKeys + values![.ellipsis, idx ..< (idx + seqLen), 0...] = newValues + offset += seqLen + idx += seqLen + + // Return current valid portion + if offset < maxSize { + return (keys![.ellipsis, .. (MLXArray, MLXArray) { + if keys.dim(2) == 1 { + return updateInPlace(keys: keys, values: values) + } + return updateConcat(keys: keys, values: values) + } + + public var isTrimmable: Bool { + offset < maxSize + } + + @discardableResult + public func trim(_ n: Int) -> Int { + let trimmed = min(offset, n) + offset -= trimmed + idx -= trimmed + return trimmed + } + + public func makeMask( + queryLength n: Int, + windowSize: Int? = nil, + returnArray: Bool = false + ) -> MLXFast.ScaledDotProductAttentionMaskMode { + if n > 1 { + let effectiveWindowSize = windowSize ?? maxSize + let effectiveOffset = min(maxSize - 1, offset) + if effectiveOffset + n > effectiveWindowSize || returnArray { + return .array(createCausalMask(n: n, offset: effectiveOffset, windowSize: effectiveWindowSize)) + } else { + return .causal + } + } else { + // Single token generation + guard let windowSize else { + return .none + } + + // May need mask when window < maxSize + if offset >= windowSize, maxSize > windowSize { + var maskIdx = idx + if maskIdx >= maxSize { + maskIdx = 0 + } + + let maskSize = offset < maxSize ? offset + 1 : maxSize + var mask = MLXArray(0 ..< Int32(maskSize)) .>= Int32(maskSize - windowSize) + mask = MLX.roll(mask, shift: maskIdx + 1) + return .array(mask) + } + return .none + } + } +} + +// MARK: - QuantizedKVCache + +/// Quantized KV cache for reduced memory usage. +/// +/// Stores keys and values in quantized format (default: 8-bit) to reduce +/// memory footprint for long context windows. +/// +/// Ported from: mlx_lm/models/cache.py::QuantizedKVCache +public final class QuantizedKVCache: KVCacheProtocol { + /// Growth step size for buffer allocation + public static let step = 256 + + /// Quantized keys: tuple of (quantized, scales, biases) + public var keys: (MLXArray, MLXArray, MLXArray?)? + + /// Quantized values: tuple of (quantized, scales, biases) + public var values: (MLXArray, MLXArray, MLXArray?)? + + public var offset: Int = 0 + + /// Quantization group size + public let groupSize: Int + + /// Bits per quantized value + public let bits: Int + + /// Creates a quantized KV cache. + /// + /// - Parameters: + /// - groupSize: Number of values per quantization group (default: 64) + /// - bits: Bits per quantized value (default: 8) + public init(groupSize: Int = 64, bits: Int = 8) { + self.groupSize = groupSize + self.bits = bits + } + + /// Returns dequantized keys and values for KV-sharing scenarios. + public var state: (keys: MLXArray, values: MLXArray)? { + guard let k = keys, let v = values, offset > 0 else { return nil } + let dequantK = MLX.dequantized(k.0, scales: k.1, biases: k.2, groupSize: groupSize, bits: bits) + let dequantV = MLX.dequantized(v.0, scales: v.1, biases: v.2, groupSize: groupSize, bits: bits) + return (dequantK[.ellipsis, .. (MLXArray, MLXArray) { + let batchSize = newKeys.dim(0) + let numKvHeads = newKeys.dim(1) + let numSteps = newKeys.dim(2) + let keyHeadDim = newKeys.dim(3) + let valueHeadDim = newValues.dim(3) + + let prev = offset + + // Calculate elements per int for this bit width + let elPerInt = 8 * MemoryLayout.size / bits + + // Check if we need to grow buffers + if keys == nil || (prev + numSteps) > keys!.0.dim(2) { + let newSteps = (Self.step + numSteps - 1) / Self.step * Self.step + let shape = [batchSize, numKvHeads, newSteps] + + func initQuant(dim: Int) -> (MLXArray, MLXArray, MLXArray?) { + ( + MLXArray.zeros(shape + [dim / elPerInt], dtype: .uint32), + MLXArray.zeros(shape + [dim / groupSize], dtype: newKeys.dtype), + MLXArray.zeros(shape + [dim / groupSize], dtype: newKeys.dtype) + ) + } + + func expandQuant(_ x: (MLXArray, MLXArray, MLXArray?)) -> (MLXArray, MLXArray, MLXArray?) { + func expand(_ arr: MLXArray) -> MLXArray { + let newArr = MLXArray.zeros(shape + [arr.dim(-1)], dtype: arr.dtype) + return concatenated([arr, newArr], axis: 2) + } + return (expand(x.0), expand(x.1), x.2.map { expand($0) }) + } + + if keys != nil { + // Trim if not aligned + if prev % Self.step != 0 { + func trimToOffset(_ x: (MLXArray, MLXArray, MLXArray?)) -> (MLXArray, MLXArray, MLXArray?) { + ( + x.0[.ellipsis, .. Int { + let trimmed = min(offset, n) + offset -= trimmed + return trimmed + } +} + +// MARK: - Factory Functions + +/// Creates prompt caches for a model. +/// +/// Defers to the model's `makeCache()` if available, otherwise creates +/// default KVCache instances for each layer. +/// +/// - Parameters: +/// - numLayers: Number of transformer layers +/// - maxKvSize: If provided, creates RotatingKVCache with this max size +/// - Returns: Array of cache instances, one per layer +public func makePromptCache( + numLayers: Int, + maxKvSize: Int? = nil +) -> [any KVCacheProtocol] { + if let maxKvSize { + (0 ..< numLayers).map { _ in + RotatingKVCache(maxSize: maxKvSize, keep: 4) + } + } else { + (0 ..< numLayers).map { _ in StandardKVCache() } + } +} + +/// Returns the maximum cache length across all caches. +public func cacheLength(_ cache: [any KVCacheProtocol]) -> Int { + cache.map(\.offset).max() ?? 0 +} + +/// Checks if all caches in the list can be trimmed. +public func canTrimPromptCache(_ cache: [any KVCacheProtocol]) -> Bool { + cache.allSatisfy(\.isTrimmable) +} + +/// Trims all caches by the specified number of tokens. +/// +/// - Returns: Actual number of tokens trimmed (from first cache) +@discardableResult +public func trimPromptCache(_ cache: [any KVCacheProtocol], numTokens: Int) -> Int { + guard canTrimPromptCache(cache), !cache.isEmpty else { return 0 } + return cache.map { $0.trim(numTokens) }.first ?? 0 +} diff --git a/packages/swift/Sources/NodeMLXCore/ported/README.md b/packages/swift/Sources/NodeMLXCore/ported/README.md new file mode 100644 index 0000000..ac7ab36 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/ported/README.md @@ -0,0 +1,90 @@ +# Ported Code + +Code in this directory is ported from Apple's `mlx-lm` Python library using LLM assistance. + +## Source + +- **Repository**: https://github.com/ml-explore/mlx-lm +- **Path**: `mlx_lm/models/` +- **Git Hash**: `7585c142a6be9c9245f4ce61d087839776cb8275` +- **Date**: 2026-01-12 + +## Ported Files + +| Python Source | Swift File | Description | +| ------------------ | -------------------- | -------------------------------------------------------- | +| `cache.py` | `KVCache.swift` | KV cache implementations (Standard, Rotating, Quantized) | +| `rope_utils.py` | `RoPEUtils.swift` | Rotary position embeddings (Standard, Llama3, Yarn, Su) | +| `switch_layers.py` | `SwitchLayers.swift` | MoE switch layers (SwitchLinear, SwitchGLU, etc.) | +| `gemma.py` | `GemmaRMSNorm.swift` | Gemma-style (1+weight) RMSNorm | + +## Porting Guidelines + +### File Header + +Every ported file must include: + +```swift +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/.py +// Git Hash: () +``` + +### Update Process + +1. **Check latest mlx-lm**: + + ```bash + curl -s "https://api.github.com/repos/ml-explore/mlx-lm/commits/main" | grep '"sha"' | head -1 + ``` + +2. **Download Python source**: + + ```bash + curl -s "https://raw.githubusercontent.com/ml-explore/mlx-lm/main/mlx_lm/models/.py" -o /tmp/.py + ``` + +3. **Use Cursor command**: `/port-python-to-swift` + +4. **Update documentation**: Update hash in file header and PORTING_DECISIONS.md + +## Design Decisions + +See [PORTING_DECISIONS.md](../../PORTING_DECISIONS.md) for detailed architectural decisions. + +### Key Patterns + +| Python | Swift | +| ------------- | ----------------- | +| `snake_case` | `camelCase` | +| `mx.array` | `MLXArray` | +| `nn.Module` | `Module` (MLXNN) | +| `__init__` | `init` | +| `@property` | computed property | +| `Optional[T]` | `T?` | + +### What We Skip + +- Batch processing (BatchKVCache, etc.) +- SSM models (MambaCache) +- Serialization (save/load prompt cache) +- Server-specific features +- Niche use cases (< 5% of users) + +## Testing + +Tests live in `packages/swift/Tests/NodeMLXCoreTests/`: + +- `KVCacheTests.swift` +- `RoPEUtilsTests.swift` +- `SwitchLayersTests.swift` + +Run tests: + +```bash +cd packages/swift +swift test +``` diff --git a/packages/swift/Sources/NodeMLXCore/ported/RoPEUtils.swift b/packages/swift/Sources/NodeMLXCore/ported/RoPEUtils.swift new file mode 100644 index 0000000..62625d8 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/ported/RoPEUtils.swift @@ -0,0 +1,384 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/rope_utils.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - RoPE Provider Protocol + +/// Protocol for all RoPE (Rotary Position Embedding) implementations. +/// +/// RoPE variants provide position information by rotating query/key vectors +/// at different frequencies depending on their position in the sequence. +public protocol RoPEProvider { + /// Applies rotary position embedding to the input tensor. + /// + /// - Parameters: + /// - x: Input tensor of shape [B, H, S, D] + /// - offset: Position offset for cached sequence + /// - Returns: Tensor with rotary embeddings applied + func callAsFunction(_ x: MLXArray, offset: Int) -> MLXArray +} + +// MARK: - Su Scaled RoPE (longrope) + +/// Su Scaled Rotary Position Embedding for extended context. +/// +/// Uses scaling factors to extend the effective context length beyond +/// the original training length. Primarily used for "longrope" models. +/// +/// Ported from: mlx_lm/models/rope_utils.py::SuScaledRoPE +public final class SuScaledRoPE: Module, RoPEProvider { + private let dim: Int + private let freqs: MLXArray + private let scale: Float + + /// Creates a Su-scaled RoPE layer. + /// + /// - Parameters: + /// - dims: Feature dimensions to rotate + /// - base: Base frequency (default: 10000) + /// - maxPositionEmbeddings: Extended context length (default: 131072) + /// - originalMaxPositionEmbeddings: Original training length (default: 4096) + /// - longFactor: Scaling factors for extended positions + /// - longMscale: Optional explicit magnitude scale + public init( + dims: Int, + base: Float = 10000.0, + maxPositionEmbeddings: Int = 131_072, + originalMaxPositionEmbeddings: Int = 4096, + longFactor: [Float] = [1.0], + longMscale: Float? = nil + ) { + dim = dims + + // Compute base frequencies + let indices = MLXArray(stride(from: Float(0), to: Float(dims), by: 2)) + let baseFreqs = pow(Float(base), indices / Float(dims)) + + // Apply long scaling factors + let factors = MLXArray(longFactor) + freqs = factors * baseFreqs + + // Compute magnitude scale + let factor = Float(maxPositionEmbeddings) / Float(originalMaxPositionEmbeddings) + if let mscale = longMscale { + scale = mscale + } else if factor <= 1.0 { + scale = 1.0 + } else { + // Default scale: sqrt(1 + log(factor) / log(original)) + scale = sqrt(1.0 + log(factor) / log(Float(originalMaxPositionEmbeddings))) + } + + super.init() + } + + public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { + // Scale the rotated dimensions + let result = x + result[.ellipsis, .. lowFreqWavelen, baseFreqs * factor, baseFreqs) + + // Medium frequencies get smooth interpolation + let isMediumFreq = logicalAnd(wavelens .> highFreqWavelen, wavelens .< lowFreqWavelen) + let smoothFactors = (Float(oldContextLen) / wavelens - lowFreqFactor) / (highFreqFactor - lowFreqFactor) + let smoothFreqs = baseFreqs / ((1.0 - smoothFactors) / factor + smoothFactors) + + freqs = which(isMediumFreq, smoothFreqs, baseFreqs) + + super.init() + } + + public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { + MLXFast.RoPE( + x, + dimensions: dims, + traditional: traditional, + base: nil, + scale: 1.0, + offset: offset, + freqs: freqs + ) + } +} + +// MARK: - Yarn RoPE + +/// Yet Another RoPE for extended context windows. +/// +/// Uses a more sophisticated frequency interpolation scheme with +/// configurable beta parameters and magnitude scaling. +/// +/// Ported from: mlx_lm/models/rope_utils.py::YarnRoPE +public final class YarnRoPE: Module, RoPEProvider { + private let dims: Int + private let traditional: Bool + private let freqs: MLXArray + private let mscale: Float + + /// Creates a YARN RoPE layer. + /// + /// - Parameters: + /// - dims: Feature dimensions to rotate + /// - traditional: Use traditional RoPE formulation + /// - maxPositionEmbeddings: Maximum sequence length + /// - base: Base frequency + /// - scalingFactor: Context extension factor + /// - originalMaxPositionEmbeddings: Original training length + /// - betaFast: High frequency correction parameter + /// - betaSlow: Low frequency correction parameter + /// - mscale: Magnitude scaling factor + /// - mscaleAllDim: Dimension-wide magnitude scaling + public init( + dims: Int, + traditional: Bool = false, + maxPositionEmbeddings _: Int = 2048, + base: Float = 10000.0, + scalingFactor: Float = 1.0, + originalMaxPositionEmbeddings: Int = 4096, + betaFast: Float = 32.0, + betaSlow: Float = 1.0, + mscale: Float = 1.0, + mscaleAllDim: Float = 0.0 + ) { + self.dims = dims + self.traditional = traditional + + // Helper functions + func yarnFindCorrectionDim(_ numRotations: Float) -> Float { + Float(dims) * log(Float(originalMaxPositionEmbeddings) / (numRotations * 2.0 * Float.pi)) / (2.0 * log(base)) + } + + func yarnFindCorrectionRange() -> (Int, Int) { + let low = Int(floor(yarnFindCorrectionDim(betaFast))) + let high = Int(ceil(yarnFindCorrectionDim(betaSlow))) + return (max(low, 0), min(high, dims - 1)) + } + + func yarnGetMscale(scale: Float, m: Float) -> Float { + if scale <= 1.0 { + return 1.0 + } + return 0.1 * m * log(scale) + 1.0 + } + + func yarnLinearRampMask(minVal: Float, maxVal: Float, dim: Int) -> MLXArray { + var maxV = maxVal + if minVal == maxVal { + maxV += 0.001 // Prevent singularity + } + let indices = MLXArray(0 ..< Int32(dim)).asType(.float32) + let linearFunc = (indices - minVal) / (maxV - minVal) + return clip(linearFunc, min: 0, max: 1) + } + + // Compute mscale + self.mscale = yarnGetMscale(scale: scalingFactor, m: mscale) / yarnGetMscale(scale: scalingFactor, m: mscaleAllDim) + + // Compute frequencies + let indices = MLXArray(stride(from: Float(0), to: Float(dims), by: 2)) + let freqExtra = pow(Float(base), indices / Float(dims)) + let freqInter = scalingFactor * freqExtra + + let (low, high) = yarnFindCorrectionRange() + let freqMask = 1.0 - yarnLinearRampMask(minVal: Float(low), maxVal: Float(high), dim: dims / 2) + + freqs = (freqInter * freqExtra) / (freqInter * freqMask + freqExtra * (1.0 - freqMask)) + + super.init() + } + + public func callAsFunction(_ x: MLXArray, offset: Int = 0) -> MLXArray { + let result = x + if mscale != 1.0 { + result[.ellipsis, .. MLXArray { + rope(x, offset: offset) + } +} + +// MARK: - Factory Function + +/// Initializes the appropriate RoPE implementation based on configuration. +/// +/// Supported rope_type values: +/// - "default": Standard RoPE +/// - "linear": Linearly scaled (scale = 1/factor) +/// - "llama3": Llama 3 with smooth frequency interpolation +/// - "yarn": YARN for extended context +/// - "longrope": Su-scaled for very long context +/// - "mrope": Multimodal (returns standard RoPE) +/// +/// - Parameters: +/// - dims: Feature dimensions to rotate +/// - base: Base frequency +/// - traditional: Use traditional RoPE formulation +/// - scalingConfig: Optional configuration dictionary +/// - maxPositionEmbeddings: Maximum sequence length +/// - Returns: Configured RoPE implementation +public func initializeRope( + dims: Int, + base: Float, + traditional: Bool, + scalingConfig: [String: Any]? = nil, + maxPositionEmbeddings: Int? = nil +) -> any RoPEProvider { + let ropeType: String = if let config = scalingConfig { + (config["type"] as? String) ?? (config["rope_type"] as? String) ?? "default" + } else { + "default" + } + + switch ropeType { + case "default": + return StandardRoPE(dims: dims, traditional: traditional, base: base, scale: 1.0) + + case "linear": + let factor = (scalingConfig?["factor"] as? Double).map { Float($0) } ?? 1.0 + return StandardRoPE(dims: dims, traditional: traditional, base: base, scale: 1.0 / factor) + + case "llama3": + return Llama3RoPE( + dims: dims, + maxPositionEmbeddings: maxPositionEmbeddings ?? 2048, + traditional: traditional, + base: base, + scalingConfig: scalingConfig ?? [:] + ) + + case "yarn": + let factor = (scalingConfig?["factor"] as? Double).map { Float($0) } ?? 1.0 + let origMax = (scalingConfig?["original_max_position_embeddings"] as? Int) ?? 4096 + let betaFast = (scalingConfig?["beta_fast"] as? Double).map { Float($0) } ?? 32.0 + let betaSlow = (scalingConfig?["beta_slow"] as? Double).map { Float($0) } ?? 1.0 + let mscale = (scalingConfig?["mscale"] as? Double).map { Float($0) } ?? 1.0 + let mscaleAllDim = (scalingConfig?["mscale_all_dim"] as? Double).map { Float($0) } ?? 0.0 + + return YarnRoPE( + dims: dims, + traditional: traditional, + maxPositionEmbeddings: maxPositionEmbeddings ?? 2048, + base: base, + scalingFactor: factor, + originalMaxPositionEmbeddings: origMax, + betaFast: betaFast, + betaSlow: betaSlow, + mscale: mscale, + mscaleAllDim: mscaleAllDim + ) + + case "longrope": + guard let config = scalingConfig else { + fatalError("longrope requires scaling configuration") + } + let origMax = config["original_max_position_embeddings"] as? Int ?? 4096 + let longFactor = config["long_factor"] as? [Double] ?? [1.0] + + return SuScaledRoPE( + dims: dims, + base: base, + maxPositionEmbeddings: maxPositionEmbeddings ?? 131_072, + originalMaxPositionEmbeddings: origMax, + longFactor: longFactor.map { Float($0) } + ) + + case "mrope": + // MRoPE: multimodal, position handling in attention + return StandardRoPE(dims: dims, traditional: traditional, base: base) + + default: + fatalError("Unsupported RoPE type: \(ropeType)") + } +} diff --git a/packages/swift/Sources/NodeMLXCore/ported/SwitchLayers.swift b/packages/swift/Sources/NodeMLXCore/ported/SwitchLayers.swift new file mode 100644 index 0000000..3851832 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/ported/SwitchLayers.swift @@ -0,0 +1,407 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/models/switch_layers.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import Foundation +import MLX +import MLXNN + +// MARK: - Helper Functions + +/// Sorts tokens by expert assignment for efficient batched access. +/// +/// When processing many tokens, sorting by expert index improves memory +/// access patterns during the expert computation. +/// +/// - Parameters: +/// - x: Input tensor [N, ...] +/// - indices: Expert indices [N, K] +/// - Returns: Tuple of (sorted x, sorted indices, inverse order for unsorting) +public func gatherSort(_ x: MLXArray, _ indices: MLXArray) -> (MLXArray, MLXArray, MLXArray) { + let m = indices.shape.last! + let flatIndices = indices.flattened() + let order = argSort(flatIndices) + let invOrder = argSort(order) + + let sortedIndices = flatIndices[order] + let sortedX = x.flattened(start: 0, end: -3)[order / m] + + return (sortedX, sortedIndices, invOrder) +} + +/// Restores original token order after expert processing. +/// +/// - Parameters: +/// - x: Sorted tensor +/// - invOrder: Inverse permutation from gatherSort +/// - shape: Optional original shape to restore +/// - Returns: Tensor in original token order +public func scatterUnsort(_ x: MLXArray, _ invOrder: MLXArray, shape: [Int]? = nil) -> MLXArray { + var result = x[invOrder] + if let shape { + result = result.reshaped([shape[0], shape[1]] + Array(result.shape.dropFirst())) + } + return result +} + +// MARK: - SwitchLinear + +/// Expert-specific linear layer for Mixture of Experts. +/// +/// Maintains separate weight matrices for each expert and uses +/// `gather_mm` for efficient batched computation. +/// +/// Ported from: mlx_lm/models/switch_layers.py::SwitchLinear +public class SwitchLinear: Module { + @ModuleInfo(key: "weight") var weight: MLXArray + @ModuleInfo(key: "bias") var bias: MLXArray? + + public var inputDims: Int { weight.dim(2) } + public var outputDims: Int { weight.dim(1) } + public var numExperts: Int { weight.dim(0) } + + /// Creates a SwitchLinear layer. + /// + /// - Parameters: + /// - inputDims: Input feature dimension + /// - outputDims: Output feature dimension + /// - numExperts: Number of expert weight matrices + /// - bias: Whether to include bias terms + public init(inputDims: Int, outputDims: Int, numExperts: Int, bias: Bool = true) { + let scale = sqrt(1.0 / Float(inputDims)) + _weight.wrappedValue = MLXRandom.uniform( + low: -scale, + high: scale, + [numExperts, outputDims, inputDims] + ) + + if bias { + _bias.wrappedValue = MLXArray.zeros([numExperts, outputDims]) + } + } + + /// Forward pass with expert selection. + /// + /// - Parameters: + /// - x: Input tensor + /// - indices: Expert indices for each token + /// - sortedIndices: Whether indices are pre-sorted + /// - Returns: Expert-weighted output + public func callAsFunction(_ x: MLXArray, indices: MLXArray, sortedIndices: Bool = false) -> MLXArray { + var result = MLX.gatherMatmul( + x, + weight.swappedAxes(-1, -2), + rhsIndices: indices, + sortedIndices: sortedIndices + ) + + if let bias { + result = result + expandedDimensions(bias[indices], axis: -2) + } + + return result + } + + /// Converts to quantized version. + public func toQuantized(groupSize: Int = 64, bits: Int = 4, mode: QuantizationMode = .affine) -> QuantizedSwitchLinear { + QuantizedSwitchLinear(self, groupSize: groupSize, bits: bits, mode: mode) + } +} + +// MARK: - QuantizedSwitchLinear + +/// Quantized version of SwitchLinear for reduced memory usage. +/// +/// Uses quantized weights with per-group scales and biases for +/// memory-efficient expert computation. +/// +/// Ported from: mlx_lm/models/switch_layers.py::QuantizedSwitchLinear +public class QuantizedSwitchLinear: Module { + @ModuleInfo(key: "weight") var weight: MLXArray + @ModuleInfo(key: "scales") var scales: MLXArray + @ModuleInfo(key: "biases") var biases: MLXArray? + @ModuleInfo(key: "bias") var bias: MLXArray? + + public let inputDims: Int + public let outputDims: Int + public let numExperts: Int + public let groupSize: Int + public let bits: Int + public let mode: QuantizationMode + + /// Creates a QuantizedSwitchLinear from an existing SwitchLinear. + public init(_ other: SwitchLinear, groupSize: Int = 64, bits: Int = 4, mode: QuantizationMode = .affine) { + inputDims = other.inputDims + outputDims = other.outputDims + numExperts = other.numExperts + self.groupSize = groupSize + self.bits = bits + self.mode = mode + + let (qw, sc, bi) = MLX.quantized(other.weight, groupSize: groupSize, bits: bits) + _weight.wrappedValue = qw + _scales.wrappedValue = sc + _biases.wrappedValue = bi + + if let otherBias = other.bias { + _bias.wrappedValue = otherBias + } + + super.init() + + // Freeze quantized weights + freeze() + } + + /// Creates a QuantizedSwitchLinear with explicit parameters. + public init( + inputDims: Int, + outputDims: Int, + numExperts: Int, + bias: Bool = true, + groupSize: Int = 64, + bits: Int = 4, + mode: QuantizationMode = .affine + ) { + self.inputDims = inputDims + self.outputDims = outputDims + self.numExperts = numExperts + self.groupSize = groupSize + self.bits = bits + self.mode = mode + + let scale = sqrt(1.0 / Float(inputDims)) + let initialWeight = MLXRandom.uniform( + low: -scale, + high: scale, + [numExperts, outputDims, inputDims] + ) + + let (qw, sc, bi) = MLX.quantized(initialWeight, groupSize: groupSize, bits: bits) + _weight.wrappedValue = qw + _scales.wrappedValue = sc + _biases.wrappedValue = bi + + if bias { + _bias.wrappedValue = MLXArray.zeros([numExperts, outputDims]) + } + + super.init() + freeze() + } + + public func callAsFunction(_ x: MLXArray, indices: MLXArray, sortedIndices: Bool = false) -> MLXArray { + var result = MLX.gatherQuantizedMatmul( + x, + weight, + scales: scales, + biases: biases, + rhsIndices: indices, + transpose: true, + groupSize: groupSize, + bits: bits, + sortedIndices: sortedIndices + ) + + if let bias { + result = result + expandedDimensions(bias[indices], axis: -2) + } + + return result + } +} + +// MARK: - SwiGLU Activation + +/// Compiled SwiGLU activation for optimal performance. +private let compiledSwiGLU: (MLXArray, MLXArray) -> MLXArray = { x, gate in + silu(gate) * x +} + +/// SwiGLU activation: SiLU(gate) * x +public func swiGLU(_ x: MLXArray, gate: MLXArray) -> MLXArray { + compiledSwiGLU(x, gate) +} + +// MARK: - SwitchGLU + +/// Gated Linear Unit with expert switching for MoE. +/// +/// Implements the standard GLU pattern with separate experts: +/// output = down_proj(activation(up_proj(x), gate_proj(x))) +/// +/// Ported from: mlx_lm/models/switch_layers.py::SwitchGLU +public class SwitchGLU: Module { + @ModuleInfo(key: "gate_proj") var gateProj: SwitchLinear + @ModuleInfo(key: "up_proj") var upProj: SwitchLinear + @ModuleInfo(key: "down_proj") var downProj: SwitchLinear + + /// Creates a SwitchGLU layer. + /// + /// - Parameters: + /// - inputDims: Input/output feature dimension + /// - hiddenDims: Hidden layer dimension + /// - numExperts: Number of experts + /// - bias: Whether to include bias terms + public init(inputDims: Int, hiddenDims: Int, numExperts: Int, bias: Bool = false) { + _gateProj.wrappedValue = SwitchLinear(inputDims: inputDims, outputDims: hiddenDims, numExperts: numExperts, bias: bias) + _upProj.wrappedValue = SwitchLinear(inputDims: inputDims, outputDims: hiddenDims, numExperts: numExperts, bias: bias) + _downProj.wrappedValue = SwitchLinear(inputDims: hiddenDims, outputDims: inputDims, numExperts: numExperts, bias: bias) + } + + public func callAsFunction(_ x: MLXArray, indices: MLXArray) -> MLXArray { + var input = expandedDimensions(x, axes: [-2, -3]) + + // Sort for efficient expert access when processing many tokens + let doSort = indices.size >= 64 + var idx = indices + var invOrder: MLXArray? + + if doSort { + (input, idx, invOrder) = gatherSort(input, indices) + } + + // GLU computation + let xUp = upProj(input, indices: idx, sortedIndices: doSort) + let xGate = gateProj(input, indices: idx, sortedIndices: doSort) + var result = downProj(swiGLU(xUp, gate: xGate), indices: idx, sortedIndices: doSort) + + // Restore original order + if doSort, let inv = invOrder { + result = scatterUnsort(result, inv, shape: Array(indices.shape)) + } + + return result.squeezed(axis: -2) + } +} + +// MARK: - SwitchMLP + +/// Simple MLP with expert switching for MoE. +/// +/// Implements: output = fc2(activation(fc1(x))) +/// +/// Ported from: mlx_lm/models/switch_layers.py::SwitchMLP +public class SwitchMLP: Module { + @ModuleInfo(key: "fc1") var fc1: SwitchLinear + @ModuleInfo(key: "fc2") var fc2: SwitchLinear + + private let activation: (MLXArray) -> MLXArray + + /// Creates a SwitchMLP layer. + /// + /// - Parameters: + /// - inputDims: Input/output feature dimension + /// - hiddenDims: Hidden layer dimension + /// - numExperts: Number of experts + /// - activation: Activation function (default: GELU) + /// - bias: Whether to include bias terms + public init( + inputDims: Int, + hiddenDims: Int, + numExperts: Int, + activation: @escaping (MLXArray) -> MLXArray = { geluApproximate($0) }, + bias: Bool = false + ) { + self.activation = activation + + _fc1.wrappedValue = SwitchLinear(inputDims: inputDims, outputDims: hiddenDims, numExperts: numExperts, bias: bias) + _fc2.wrappedValue = SwitchLinear(inputDims: hiddenDims, outputDims: inputDims, numExperts: numExperts, bias: bias) + } + + public func callAsFunction(_ x: MLXArray, indices: MLXArray) -> MLXArray { + var input = expandedDimensions(x, axes: [-2, -3]) + + // Sort for efficient expert access + let doSort = indices.size >= 64 + var idx = indices + var invOrder: MLXArray? + + if doSort { + (input, idx, invOrder) = gatherSort(input, indices) + } + + // MLP computation + var result = fc1(input, indices: idx, sortedIndices: doSort) + result = activation(result) + result = fc2(result, indices: idx, sortedIndices: doSort) + + // Restore original order + if doSort, let inv = invOrder { + result = scatterUnsort(result, inv, shape: Array(indices.shape)) + } + + return result.squeezed(axis: -2) + } +} + +// MARK: - GPT-OSS Specific SwiGLU + +/// GPT-OSS specific SwiGLU with clipping for numerical stability. +/// +/// Uses hard clipping on the gate value before SiLU activation. +private func gptOssSwiGLU(_ x: MLXArray, gate: MLXArray, limit: Float = 7.0) -> MLXArray { + let clippedGate = clip(gate, min: -limit, max: limit) + return silu(clippedGate) * x +} + +/// Compiled version of GPT-OSS SwiGLU for optimal performance. +private let compiledGptOssSwiGLU: (MLXArray, MLXArray) -> MLXArray = { x, gate in + gptOssSwiGLU(x, gate: gate) +} + +/// GPT-OSS variant of SwitchGLU with clipped activation. +public class SwiGLUSwitchGLU: Module { + @ModuleInfo(key: "gate_proj") var gateProj: SwitchLinear + @ModuleInfo(key: "up_proj") var upProj: SwitchLinear + @ModuleInfo(key: "down_proj") var downProj: SwitchLinear + + public init(inputDims: Int, hiddenDims: Int, numExperts: Int, bias: Bool = false) { + _gateProj.wrappedValue = SwitchLinear(inputDims: inputDims, outputDims: hiddenDims, numExperts: numExperts, bias: bias) + _upProj.wrappedValue = SwitchLinear(inputDims: inputDims, outputDims: hiddenDims, numExperts: numExperts, bias: bias) + _downProj.wrappedValue = SwitchLinear(inputDims: hiddenDims, outputDims: inputDims, numExperts: numExperts, bias: bias) + } + + public func callAsFunction(_ x: MLXArray, indices: MLXArray) -> MLXArray { + var input = expandedDimensions(x, axes: [-2, -3]) + + let doSort = indices.size >= 64 + var idx = indices + var invOrder: MLXArray? + + if doSort { + (input, idx, invOrder) = gatherSort(input, indices) + } + + let xUp = upProj(input, indices: idx, sortedIndices: doSort) + let xGate = gateProj(input, indices: idx, sortedIndices: doSort) + var result = downProj(compiledGptOssSwiGLU(xUp, xGate), indices: idx, sortedIndices: doSort) + + if doSort, let inv = invOrder { + result = scatterUnsort(result, inv, shape: Array(indices.shape)) + } + + return result.squeezed(axis: -2) + } +} + +// MARK: - Weight Conversion Utilities + +/// Converts packed MoE tensors from blocks+scales format. +/// +/// Used during model loading to transform the packed tensor format +/// used in some quantized MoE checkpoints. +/// +/// - Parameters: +/// - blocks: Quantized weight blocks +/// - scales: Quantization scales +/// - Returns: Transformed tensor suitable for weight loading +public func convertMoePackedTensors(blocks: MLXArray, scales _: MLXArray) -> MLXArray { + // Interleave scales with blocks for the expected format + // This matches the pattern from mlx-swift-lm's GPTOSS.swift + // For now, return the blocks directly (scales handled separately) + blocks +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/AltUpBlock.swift b/packages/swift/Sources/NodeMLXCore/shared/AltUpBlock.swift new file mode 100644 index 0000000..92c484a --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/AltUpBlock.swift @@ -0,0 +1,122 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// AltUp (Alternating Updates) block for efficient sparse computation. +// Used by Gemma3n and potentially future models with similar architecture. + +import Foundation +import MLX +import MLXFast +import MLXNN + +/// AltUp (Alternating Updates) module for efficient sparse computation. +/// +/// AltUp reduces computation by maintaining multiple "virtual" hidden states +/// but only computing attention/MLP on one active state at a time. +/// The predict step spreads information, and the correct step refines predictions. +/// +/// Architecture: +/// 1. **Predict**: Use learned coefficients to predict inactive states from active +/// 2. **Activate**: Run attention/MLP on the active state only +/// 3. **Correct**: Refine predictions based on the activated output +/// +/// This allows N× throughput improvement with minimal quality loss. +public class AltUpBlock: Module { + public let numInputs: Int + public let activeIdx: Int + public let hiddenSize: Int + public let altupCoefClip: Float? + + @ModuleInfo(key: "correct_output_scale") public var correctOutputScale: MLXArray + @ModuleInfo(key: "correction_coefs") public var correctionCoefs: Linear + @ModuleInfo(key: "prediction_coefs") public var predictionCoefs: Linear + @ModuleInfo(key: "modality_router") public var modalityRouter: Linear + @ModuleInfo(key: "router_norm") public var routerNorm: RMSNorm + + public init(_ config: Config) { + numInputs = config.altupNumInputs + activeIdx = config.altupActiveIdx + hiddenSize = config.hiddenSize + altupCoefClip = config.altupCoefClip + + _correctOutputScale.wrappedValue = MLXArray.zeros([config.hiddenSize]) + _correctionCoefs.wrappedValue = Linear(numInputs, numInputs, bias: false) + _predictionCoefs.wrappedValue = Linear(numInputs, numInputs * numInputs, bias: false) + _modalityRouter.wrappedValue = Linear(config.hiddenSize, numInputs, bias: false) + _routerNorm.wrappedValue = RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + } + + /// Compute router modalities from input hidden states. + public func computeRouterModalities(_ x: MLXArray) -> MLXArray { + let scale = Foundation.pow(Float(hiddenSize), -1.0) + let routerInputs = routerNorm(x) * scale + let routed = modalityRouter(routerInputs).asType(.float32) + return tanh(routed) + } + + /// Predict step: modifies input using learned coefficients. + /// + /// - Parameter hiddenStates: [numInputs, batch, seq, hidden] + /// - Returns: Predictions [numInputs, batch, seq, hidden] + public func predict(_ hiddenStates: MLXArray) -> MLXArray { + let modalities = computeRouterModalities(hiddenStates[activeIdx]) + + // Compute prediction coefficients with optional clipping + var weight = predictionCoefs.weight.asType(.float32) + if let clipVal = altupCoefClip { + weight = clip(weight, min: -clipVal, max: clipVal) + } + + // Manual linear: modalities @ weight.T + var allCoefs = matmul(modalities.asType(.float32), weight.T) + let shape = modalities.shape + allCoefs = allCoefs.reshaped([shape[0], shape[1], numInputs, numInputs]) + allCoefs = allCoefs.transposed(0, 1, 3, 2) + + // Convert to float32 for better precision + let xUp = hiddenStates.asType(.float32) + let xPermuted = xUp.transposed(1, 2, 3, 0) + var predictions = matmul(xPermuted, allCoefs) + predictions = predictions.transposed(3, 0, 1, 2) + predictions = predictions + xUp + + return predictions.asType(hiddenStates.dtype) + } + + /// Correct step: refines predictions based on activated output. + /// + /// - Parameters: + /// - predictions: Predicted states [numInputs, batch, seq, hidden] + /// - activated: Output from attention/MLP [batch, seq, hidden] + /// - Returns: Corrected states [numInputs, batch, seq, hidden] + public func correct(_ predictions: MLXArray, activated: MLXArray) -> MLXArray { + let modalities = computeRouterModalities(activated) + + // Compute correction coefficients with optional clipping + var weight = correctionCoefs.weight.asType(.float32) + if let clipVal = altupCoefClip { + weight = clip(weight, min: -clipVal, max: clipVal) + } + + // Manual linear + 1.0: modalities @ weight.T + 1.0 + var allCoefs = matmul(modalities.asType(.float32), weight.T) + 1.0 + let activeX = predictions[activeIdx] + let innovation = activated - activeX + + // allCoefs: [batch, seq, numInputs] -> [numInputs, batch, seq] + allCoefs = allCoefs.transposed(2, 0, 1) + + // innovation: [batch, seq, hidden] + // Broadcast: [numInputs, batch, seq, 1] * [1, batch, seq, hidden] + let innovationExpanded = innovation.expandedDimensions(axis: 0) + let allCoefsExpanded = allCoefs.expandedDimensions(axis: -1) + let corrected = innovationExpanded * allCoefsExpanded + predictions + + return corrected.asType(activated.dtype) + } + + /// Scale the correction output (used when altupCorrectScale is enabled). + public func scaleCorrectOutput(_ corrected: MLXArray) -> MLXArray { + corrected * correctOutputScale + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/FusedQKVAttention.swift b/packages/swift/Sources/NodeMLXCore/shared/FusedQKVAttention.swift new file mode 100644 index 0000000..aba7343 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/FusedQKVAttention.swift @@ -0,0 +1,91 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Generic Fused QKV Attention layer for models using qkv_proj. +// Used by Phi3, Phi4, and similar architectures. + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Fused QKV Attention + +/// Attention layer using a single fused qkv_proj projection. +/// +/// This is more efficient than separate q/k/v projections as it requires +/// only one matrix multiplication instead of three. +/// +/// Usage in generated models: +/// ```swift +/// typealias Phi3Attention = FusedQKVAttention +/// ``` +public class FusedQKVAttention: Module { + @ModuleInfo(key: "qkv_proj") var qkvProj: Linear + @ModuleInfo(key: "o_proj") var oProj: Linear + + public let numHeads: Int + public let numKVHeads: Int + public let headDim: Int + public let scale: Float + public let rope: RoPE + + public init(_ config: C) { + numHeads = config.numAttentionHeads + numKVHeads = config.numKeyValueHeads + headDim = config.headDim + scale = config.attentionScale ?? (1.0 / sqrt(Float(headDim))) + + let qDim = numHeads * headDim + let kvDim = numKVHeads * headDim + let opSize = qDim + 2 * kvDim + + _qkvProj.wrappedValue = Linear(config.hiddenSize, opSize, bias: false) + _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: false) + rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) + } + + public func callAsFunction( + _ hiddenStates: MLXArray, + mask: MLXFast.ScaledDotProductAttentionMaskMode, + cache: inout KVCacheProtocol? + ) -> MLXArray { + let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) + + let qkv = qkvProj(hiddenStates) + let queryPos = numHeads * headDim + let kvPos = queryPos + numKVHeads * headDim + + var queries = qkv[0..., 0..., .. [B, L, hidden] + let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) + return oProj(outputReshaped) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/LaurelBlock.swift b/packages/swift/Sources/NodeMLXCore/shared/LaurelBlock.swift new file mode 100644 index 0000000..6a5f9b5 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/LaurelBlock.swift @@ -0,0 +1,47 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Laurel (Learned Augmented Residual) block for efficient low-rank residual computation. +// Used by Gemma3n and potentially future models with similar architecture. + +import MLX +import MLXNN + +/// Laurel (Learned Augmented Residual) block. +/// +/// A low-rank residual layer that adds a learned residual to the input: +/// output = x + postNorm(right(left(x))) +/// +/// This is more parameter-efficient than full-rank residual connections +/// while still allowing the model to learn useful residual transformations. +/// +/// Architecture: +/// 1. Project down to low-rank: x → Linear(hidden → laurelRank) +/// 2. Project back up: Linear(laurelRank → hidden) +/// 3. Normalize: RMSNorm +/// 4. Add residual: x + normalized +public class LaurelBlock: Module { + @ModuleInfo(key: "linear_left") public var linearLeft: Linear + @ModuleInfo(key: "linear_right") public var linearRight: Linear + @ModuleInfo(key: "post_laurel_norm") public var postLaurelNorm: RMSNorm + + private let hiddenSize: Int + private let laurelRank: Int + + public init(_ config: Config) { + hiddenSize = config.hiddenSize + laurelRank = config.laurelRank + + _linearLeft.wrappedValue = Linear(hiddenSize, laurelRank, bias: false) + _linearRight.wrappedValue = Linear(laurelRank, hiddenSize, bias: false) + _postLaurelNorm.wrappedValue = RMSNorm(dimensions: hiddenSize, eps: config.rmsNormEps) + } + + public func callAsFunction(_ x: MLXArray) -> MLXArray { + var laurel = linearLeft(x) + laurel = linearRight(laurel) + laurel = postLaurelNorm(laurel) + // Add residual connection + return x + laurel + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/MathUtils.swift b/packages/swift/Sources/NodeMLXCore/shared/MathUtils.swift new file mode 100644 index 0000000..934181d --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/MathUtils.swift @@ -0,0 +1,65 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Mathematical utility functions for neural network operations. + +import Foundation +import MLX + +// MARK: - Math Utilities + +/// Mathematical utility functions used across the codebase. +public enum MathUtils { + /// Approximate inverse error function. + /// + /// Uses a rational approximation that is accurate to about 4 decimal places. + /// Used primarily for computing gelu_topk sparse activation thresholds. + /// + /// - Parameter x: Input value in range (-1, 1) + /// - Returns: Inverse error function of x + public static func erfinv(_ x: Float) -> Float { + let a: Float = 0.147 + let sign: Float = x < 0 ? -1 : 1 + let x2 = x * x + let lnTerm = log(1 - x2) + let term1 = 2 / (Float.pi * a) + lnTerm / 2 + let term2 = lnTerm / a + return sign * sqrt(sqrt(term1 * term1 - term2) - term1) + } + + /// Clip residual for float16 overflow protection. + /// + /// When using float16, residual additions can overflow. This function + /// converts to float32 for the addition and clips to float16 bounds + /// before converting back. + /// + /// - Parameters: + /// - x: First operand + /// - y: Second operand to add + /// - Returns: Clipped sum in original dtype + public static func clipResidual(_ x: MLXArray, _ y: MLXArray) -> MLXArray { + if x.dtype != .float16 { + return x + y + } + let bound = Float16.greatestFiniteMagnitude + let sum = (x.asType(.float32) + y.asType(.float32)) + return clip(sum, min: MLXArray(-Float(bound)), max: MLXArray(Float(bound))).asType(.float16) + } + + /// Top-k selection for MoE routing. + /// + /// Efficiently selects the top k values and their indices from an array. + /// Uses argPartition for O(n) performance instead of O(n log n) full sort. + /// + /// - Parameters: + /// - a: Input array + /// - k: Number of top elements to select + /// - axis: Axis along which to select (default: -1) + /// - Returns: Tuple of (top k values, top k indices) + public static func topK(_ a: MLXArray, k: Int, axis: Int = -1) -> (values: MLXArray, indices: MLXArray) { + let partitionedIndices = argPartition(a, kth: -k, axis: axis) + let topKIndices = partitionedIndices[.ellipsis, (-k)...] + let topKValues = takeAlong(a, topKIndices, axis: axis) + return (topKValues, topKIndices) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/MoESanitizer.swift b/packages/swift/Sources/NodeMLXCore/shared/MoESanitizer.swift new file mode 100644 index 0000000..0a235aa --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/MoESanitizer.swift @@ -0,0 +1,117 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// MoE (Mixture of Experts) weight sanitization utilities. +// Used by GPT-OSS and similar MoE architectures. + +import Foundation +import MLX +import MLXNN + +// MARK: - MoE Weight Sanitizer + +/// Utilities for sanitizing MoE model weights. +/// +/// Handles: +/// - Packed tensor format (blocks + scales) → unpacked bfloat16 +/// - Fused gate_up_proj → separate gate_proj + up_proj +/// - Weight key transformations for MLXNN compatibility +public enum MoESanitizer { + /// Convert packed MoE tensors from blocks+scales format to unpacked bfloat16. + /// + /// The packed format uses a 4-bit lookup table encoding with separate scale factors. + /// This function unpacks them into standard bfloat16 weights. + /// + /// - Parameters: + /// - blocks: Packed weight blocks + /// - scales: Scale factors for each block + /// - Returns: Unpacked weights in bfloat16 format + public static func convertPackedTensors(blocks: MLXArray, scales: MLXArray) -> MLXArray { + precondition( + blocks.shape.dropLast() == scales.shape, + "blocks.shape=\(blocks.shape) does not match scales.shape=\(scales.shape)" + ) + + var scales = scales.asType(.int32) - 127 + let lut = MLXArray([ + +0.0, +0.5, +1.0, +1.5, +2.0, +3.0, +4.0, +6.0, + -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0, + ]).asType(.bfloat16) + + let (prefixShape, G, B) = (Array(blocks.shape.dropLast(2)), blocks.dim(-2), blocks.dim(-1)) + + let blocks = blocks.reshaped(-1, B) + scales = scales.reshaped(-1, 1) + + let idxLo = blocks & 0x0F + let idxHi = blocks >> 4 + + var out = stacked([lut[idxLo], lut[idxHi]], axis: -1).flattened(start: -2) + out = (2.0 ** scales) * out + out = out.reshaped(prefixShape + [G * B * 2]) + return out.asType(.bfloat16) + } + + /// Sanitize MoE model weights for MLXNN compatibility. + /// + /// Performs the following transformations: + /// 1. Unpacks packed tensors (blocks + scales) if present + /// 2. Splits fused gate_up_proj into separate gate_proj and up_proj + /// 3. Transforms weight keys to match MLXNN module expectations + /// + /// - Parameter weights: Raw weights dictionary from model file + /// - Returns: Sanitized weights ready for MLXNN module loading + public static func sanitize(weights: [String: MLXArray]) -> [String: MLXArray] { + var weights = weights + + // Check if already in expected format + if weights.keys.contains(where: { $0.contains("gate_proj.weight") }) { + return weights + } + + // Handle packed MoE tensor format (blocks + scales) + if weights.keys.contains(where: { $0.contains("gate_up_proj_scales") }) { + var newWeights: [String: MLXArray] = [:] + for (k, v) in weights { + if k.hasSuffix("_scales") { + continue + } else if k.hasSuffix("_blocks") { + let scaleKey = k.replacingOccurrences(of: "_blocks", with: "_scales") + if let scales = weights[scaleKey] { + let newV = convertPackedTensors(blocks: v, scales: scales) + let newK = k.replacingOccurrences(of: "_blocks", with: "") + newWeights[newK] = newV + } + } else { + newWeights[k] = v + } + } + weights = newWeights + } + + // Transform weight keys to expected format + var finalWeights: [String: MLXArray] = [:] + for (k, v) in weights { + if k.contains("gate_up_proj"), !k.contains("bias") { + // Split interleaved gate_up_proj into separate gate_proj and up_proj + finalWeights[k.replacingOccurrences(of: "gate_up_proj", with: "gate_proj.weight")] = + v[.ellipsis, .stride(by: 2), 0...] + finalWeights[k.replacingOccurrences(of: "gate_up_proj", with: "up_proj.weight")] = + v[.ellipsis, .stride(from: 1, by: 2), 0...] + } else if k.contains("down_proj"), !k.contains("bias") { + finalWeights[k.replacingOccurrences(of: "down_proj", with: "down_proj.weight")] = v + } else if k.contains("gate_up_proj_bias") { + finalWeights[k.replacingOccurrences(of: "gate_up_proj_bias", with: "gate_proj.bias")] = + v[.ellipsis, .stride(by: 2)] + finalWeights[k.replacingOccurrences(of: "gate_up_proj_bias", with: "up_proj.bias")] = + v[.ellipsis, .stride(from: 1, by: 2)] + } else if k.contains("down_proj_bias") { + finalWeights[k.replacingOccurrences(of: "down_proj_bias", with: "down_proj.bias")] = v + } else { + finalWeights[k] = v + } + } + + return finalWeights + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/Protocols.swift b/packages/swift/Sources/NodeMLXCore/shared/Protocols.swift new file mode 100644 index 0000000..da7e0f7 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/Protocols.swift @@ -0,0 +1,145 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Shared protocols for LLM model configurations. +// These enable generic implementations of common model components. + +import Foundation +import MLX + +// MARK: - Base Configuration Protocol + +/// Common configuration properties shared by all transformer models. +public protocol BaseModelConfiguration: Decodable, Sendable { + var hiddenSize: Int { get } + var numHiddenLayers: Int { get } + var numAttentionHeads: Int { get } + var numKeyValueHeads: Int { get } + var intermediateSize: Int { get } + var vocabSize: Int { get } + var headDim: Int { get } + var rmsNormEps: Float { get } + var ropeTheta: Float { get } + var maxPositionEmbeddings: Int { get } + var attentionBias: Bool { get } + var mlpBias: Bool { get } + var ropeScaling: [String: StringOrNumber]? { get } +} + +// MARK: - Attention Configuration + +/// Configuration for attention layers with all required parameters. +/// Models conforming to this can use the shared FusedQKVAttention or SeparateQKVAttention. +public protocol AttentionConfiguration: BaseModelConfiguration { + /// Attention scale override (nil = use 1/sqrt(headDim)) + var attentionScale: Float? { get } +} + +public extension AttentionConfiguration { + var attentionScale: Float? { nil } +} + +// MARK: - Sliding Window Configuration + +/// Configuration for models with sliding window attention (Mistral, etc.) +public protocol SlidingWindowConfiguration: BaseModelConfiguration { + var slidingWindow: Int { get } + var slidingWindowPattern: Int { get } + + /// Check if a layer uses global attention + func isGlobalLayer(_ layerIdx: Int) -> Bool +} + +public extension SlidingWindowConfiguration { + func isGlobalLayer(_ layerIdx: Int) -> Bool { + (layerIdx % slidingWindowPattern) == (slidingWindowPattern - 1) + } +} + +// MARK: - MoE Configuration + +/// Configuration for Mixture of Experts models (GPT-OSS, etc.) +public protocol MoEConfiguration: BaseModelConfiguration { + var numLocalExperts: Int { get } + var numExpertsPerTok: Int { get } +} + +// MARK: - Sparse MLP Configuration + +/// Configuration for MLP layers with sparse activation (Gemma3n). +public protocol SparseMLPConfiguration: BaseModelConfiguration { + /// Per-layer intermediate sizes + var intermediateSizes: [Int] { get } + + /// Per-layer activation sparsity pattern + var activationSparsityPattern: [Float] { get } + + /// Get intermediate size for a specific layer + func intermediateSize(forLayer idx: Int) -> Int +} + +// MARK: - AltUp Configuration + +/// Configuration for models with Alternating Updates (AltUp) architecture. +/// Used by Gemma3n and similar efficient sparse computation models. +public protocol AltUpConfiguration: BaseModelConfiguration { + /// Number of inputs to the AltUp module + var altupNumInputs: Int { get } + + /// Active index for predict/correct operations + var altupActiveIdx: Int { get } + + /// Optional coefficient clipping for numerical stability + var altupCoefClip: Float? { get } + + /// Whether to scale the correction output + var altupCorrectScale: Bool { get } +} + +// MARK: - Laurel Configuration + +/// Configuration for models with Laurel (Learned Augmented Residual) blocks. +public protocol LaurelConfiguration: BaseModelConfiguration { + /// Rank of the low-rank residual layer + var laurelRank: Int { get } +} + +// MARK: - Configuration Decoding Helper + +/// Helper struct for decoding model configurations from JSON. +/// Handles both top-level and nested text_config patterns. +public struct ConfigDecoder { + private let container: KeyedDecodingContainer + private let textConfigKey: Keys? + + public init(container: KeyedDecodingContainer, textConfigKey: Keys? = nil) { + self.container = container + self.textConfigKey = textConfigKey + } + + /// Decode a value, trying text_config first if available, then top-level. + public func decode(_ key: Keys, default defaultValue: T? = nil) throws -> T { + // Try nested text_config first + if let textKey = textConfigKey, + let nested = try? container.nestedContainer(keyedBy: Keys.self, forKey: textKey), + let value = try? nested.decode(T.self, forKey: key) + { + return value + } + + // Try top-level + if let value = try? container.decode(T.self, forKey: key) { + return value + } + + // Use default if provided + if let defaultValue { + return defaultValue + } + + throw DecodingError.keyNotFound( + key, + DecodingError.Context(codingPath: [], debugDescription: "Missing \(key)") + ) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/README.md b/packages/swift/Sources/NodeMLXCore/shared/README.md new file mode 100644 index 0000000..3a574ca --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/README.md @@ -0,0 +1,95 @@ +# Shared Components + +Reusable Swift implementations shared across all generated models. These reduce generated code by ~70% and provide a single source of truth for common patterns. + +## Protocols + +Configuration protocols enable generic components: + +| Protocol | Properties | Used By | +| ---------------------------- | ------------------------------------- | --------------------- | +| `BaseModelConfiguration` | hiddenSize, numHeads, ropeTheta, etc. | All models | +| `AttentionConfiguration` | + attentionScale | Phi (fused attention) | +| `SlidingWindowConfiguration` | + slidingWindow, isGlobalLayer() | Mistral | +| `MoEConfiguration` | + numExperts, numExpertsPerTok | GPT-OSS | +| `AltUpConfiguration` | + altupNumInputs, altupActiveIdx | Gemma3n | +| `LaurelConfiguration` | + laurelRank | Gemma3n | +| `SparseMLPConfiguration` | + intermediateSizes, sparsityPattern | Gemma3n | + +## Standard Components + +Generic implementations for common transformer patterns: + +| Component | Description | Used By | +| ------------------------- | ------------------------------ | ------------ | +| `RMSNorm` | Root Mean Square normalization | Most models | +| `StandardAttention` | GQA attention with RoPE | Llama, Qwen2 | +| `StandardMLP` | SwiGLU MLP (gate/up/down) | Llama, Qwen2 | +| `StandardDecoderLayer` | Pre-norm decoder (2 norms) | Llama, Qwen2 | +| `FusedQKVAttention` | Fused Q/K/V projection | Phi3, Phi4 | + +## Specialized Components + +For advanced architectures: + +| Component | Description | Used By | +| ----------------- | -------------------------------------- | ----------- | +| `AltUpBlock` | Alternating Updates for sparse compute | Gemma3n | +| `LaurelBlock` | Low-rank residual layer | Gemma3n | +| `SparseMLP` | gelu_topk sparse activation | Gemma3n | +| `MoESanitizer` | MoE weight transformation | GPT-OSS | +| `WeightSanitizer` | Standard weight cleanup | Most models | + +## Utilities + +| File | Functions | Purpose | +| ----------------- | -------------------------------------- | -------------------- | +| `MathUtils.swift` | `erfinv()`, `clipResidual()`, `topK()` | Mathematical helpers | +| `Protocols.swift` | `ConfigDecoder` | JSON decoding helper | + +## Usage in Generated Code + +### Simple Models (Llama, Qwen2) + +Generator produces typealiases: + +```swift +// MARK: - Attention +typealias LlamaAttention = StandardAttention + +// MARK: - MLP +typealias LlamaMLP = StandardMLP + +// MARK: - Decoder Layer +typealias LlamaDecoderLayer = StandardDecoderLayer +``` + +### Complex Models (Gemma3n, GPT-OSS) + +Generator produces custom code but still uses shared components: + +```swift +// Uses shared AltUpBlock +extension Gemma3nConfiguration: AltUpConfiguration {} +typealias Gemma3nAltUp = AltUpBlock + +// Uses shared MathUtils +private func clipResidual(_ x: MLXArray, _ y: MLXArray) -> MLXArray { + MathUtils.clipResidual(x, y) +} +``` + +## Adding New Components + +1. Create Swift file in `shared/` +2. Define protocol if configuration-dependent +3. Implement as generic class: `class MyComponent: Module` +4. Update generator to use component when features match +5. Add tests + +## Benefits + +- **Less code**: ~195 lines vs ~350 lines per simple model +- **Testable**: Components tested once, used everywhere +- **Consistent**: Same behavior across all models +- **Maintainable**: Fix once, applies to all models diff --git a/packages/swift/Sources/NodeMLXCore/shared/RMSNorm.swift b/packages/swift/Sources/NodeMLXCore/shared/RMSNorm.swift new file mode 100644 index 0000000..65006ac --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/RMSNorm.swift @@ -0,0 +1,30 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Shared RMSNorm implementation used by all transformer models. + +import MLX +import MLXFast +import MLXNN + +/// Root Mean Square Layer Normalization. +/// +/// This is the standard normalization layer used in modern LLMs like +/// Llama, Qwen, Mistral, Gemma, etc. +/// +/// RMSNorm is computationally simpler than LayerNorm as it only +/// normalizes by the RMS of activations, without centering. +public class RMSNorm: Module { + public let eps: Float + + @ModuleInfo(key: "weight") public var weight: MLXArray + + public init(dimensions: Int, eps: Float = 1e-6) { + self.eps = eps + _weight.wrappedValue = MLXArray.ones([dimensions]) + } + + public func callAsFunction(_ x: MLXArray) -> MLXArray { + MLXFast.rmsNorm(x, weight: weight, eps: eps) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/SamplingUtils.swift b/packages/swift/Sources/NodeMLXCore/shared/SamplingUtils.swift new file mode 100644 index 0000000..2978968 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/SamplingUtils.swift @@ -0,0 +1,156 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Sampling utilities for token generation. +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: mlx_lm/sample_utils.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import Foundation +import MLX + +// MARK: - Sampling Utilities + +/// Sampling utilities for nucleus (top-p), top-k, and min-p sampling. +public enum SamplingUtils { + // MARK: - Top-P (Nucleus) Sampling + + /// Applies top-p (nucleus) sampling to logits. + /// + /// Masks tokens outside the smallest set of tokens whose cumulative + /// probability exceeds p. + /// + /// - Parameters: + /// - logits: Input logits, shape [..., vocab_size] + /// - p: Probability threshold (0.0-1.0) + /// - Returns: Filtered logits with low-probability tokens masked to -inf + public static func applyTopP(_ logits: MLXArray, p: Float) -> MLXArray { + // Get probabilities + let probs = softmax(logits, axis: -1) + + // Sort probabilities descending + let sortedIndices = argSort(-probs, axis: -1) + let sortedProbs = takeAlong(probs, sortedIndices, axis: -1) + + // Cumulative probabilities + let cumProbs = cumsum(sortedProbs, axis: -1) + + // Create shifted cumsum: prepend 0 and drop last element + // This ensures we keep at least the top token even if it exceeds p + let zerosShape = Array(cumProbs.shape.dropLast()) + [1] + let zeros = MLXArray.zeros(zerosShape) + let shiftedCumProbs = concatenated([zeros, cumProbs[.ellipsis, ..<(-1)]], axis: -1) + + // Find cutoff: positions where shifted cumsum > p should be masked + let topPMask = shiftedCumProbs .> MLXArray(p) + + // Apply mask: set excluded tokens to -inf + let sortedLogits = takeAlong(logits, sortedIndices, axis: -1) + let filteredSortedLogits = which(topPMask, MLXArray(-Float.infinity), sortedLogits) + + // Unsort back to original order + let unsortIndices = argSort(sortedIndices, axis: -1) + return takeAlong(filteredSortedLogits, unsortIndices, axis: -1) + } + + // MARK: - Top-K Sampling + + /// Applies top-k sampling to logits. + /// + /// Keeps only the k tokens with highest probability, masking the rest. + /// + /// - Parameters: + /// - logits: Input logits, shape [..., vocab_size] + /// - k: Number of tokens to keep + /// - Returns: Filtered logits with low-probability tokens masked to -inf + public static func applyTopK(_ logits: MLXArray, k: Int) -> MLXArray { + guard k > 0 else { return logits } + + // Get top k indices using partition (more efficient than full sort) + let topKIndices = argPartition(-logits, kth: k, axis: -1)[.ellipsis, .. MLXArray { + guard minP > 0 else { return logits } + + // Get probabilities + let probs = softmax(logits, axis: -1) + + // Find maximum probability + let maxProb = probs.max(axis: -1, keepDims: true) + + // Threshold is minP * maxProb + let threshold = maxProb * MLXArray(minP) + + // Mask tokens below threshold + let mask = probs .< threshold + return which(mask, MLXArray(-Float.infinity), logits) + } + + // MARK: - Combined Sampling + + /// Samples a token from logits with temperature and optional filtering. + /// + /// - Parameters: + /// - logits: Input logits, shape [vocab_size] or [1, vocab_size] + /// - temperature: Temperature for scaling (0 = greedy) + /// - topP: Top-p threshold (1.0 = disabled) + /// - topK: Top-k count (0 = disabled) + /// - minP: Min-p threshold (0.0 = disabled) + /// - Returns: Sampled token index + public static func sampleToken( + logits: MLXArray, + temperature: Float = 1.0, + topP: Float = 1.0, + topK: Int = 0, + minP: Float = 0.0 + ) -> Int { + // Ensure 2D shape + var workingLogits = logits.ndim == 1 ? logits.reshaped([1, -1]) : logits + + // Greedy decoding + if temperature == 0 { + return argMax(workingLogits, axis: -1).item(Int.self) + } + + // Apply temperature + workingLogits = workingLogits / MLXArray(temperature) + + // Apply filters in order + if topK > 0 { + workingLogits = applyTopK(workingLogits, k: topK) + } + if topP < 1.0 { + workingLogits = applyTopP(workingLogits, p: topP) + } + if minP > 0 { + workingLogits = applyMinP(workingLogits, minP: minP) + } + + // Sample from distribution + let probs = softmax(workingLogits, axis: -1) + return categorical(probs.squeezed()).item(Int.self) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/SparseMLP.swift b/packages/swift/Sources/NodeMLXCore/shared/SparseMLP.swift new file mode 100644 index 0000000..bbd2185 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/SparseMLP.swift @@ -0,0 +1,69 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Generic Sparse MLP layer with gelu_topk activation. +// Used by Gemma3n and similar architectures. + +import Foundation +import MLX +import MLXFast +import MLXNN + +// MARK: - Sparse MLP + +/// MLP layer with optional sparse gelu_topk activation. +/// +/// The gelu_topk activation zeros out activations below a dynamic threshold +/// computed from the input statistics, enabling more efficient sparse computation. +/// +/// Usage in generated models: +/// ```swift +/// typealias Gemma3nMLP = SparseMLP +/// ``` +public class SparseMLP: Module { + @ModuleInfo(key: "gate_proj") var gateProj: Linear + @ModuleInfo(key: "up_proj") var upProj: Linear + @ModuleInfo(key: "down_proj") var downProj: Linear + + public let activationSparsity: Float + public let stdMultiplier: Float? + + public init(_ config: C, layerIdx: Int = 0) { + let intermediateSize = config.intermediateSize(forLayer: layerIdx) + _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: false) + _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: false) + _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: false) + + // Get activation sparsity for this layer + if layerIdx < config.activationSparsityPattern.count { + activationSparsity = config.activationSparsityPattern[layerIdx] + } else { + activationSparsity = 0.0 + } + + // Precompute std multiplier for gelu_topk if sparsity > 0 + if activationSparsity > 0 { + // sqrt(2) * erfinv(2 * sparsity - 1) + stdMultiplier = Float(sqrt(2.0)) * MathUtils.erfinv(2.0 * activationSparsity - 1.0) + } else { + stdMultiplier = nil + } + } + + public func callAsFunction(_ x: MLXArray) -> MLXArray { + let gateOutput = gateProj(x) + let activations: MLXArray + + if let stdMult = stdMultiplier, activationSparsity > 0 { + // gelu_topk: sparse activation + let inputMean = mean(gateOutput, axis: -1, keepDims: true) + let inputStd = sqrt(mean((gateOutput - inputMean).pow(2), axis: -1, keepDims: true)) + let cutoffX = inputMean + inputStd * stdMult + activations = geluApproximate(maximum(MLXArray(Float(0)), gateOutput - cutoffX)) + } else { + activations = geluApproximate(gateOutput) + } + + return downProj(activations * upProj(x)) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/StandardAttention.swift b/packages/swift/Sources/NodeMLXCore/shared/StandardAttention.swift new file mode 100644 index 0000000..0a04a72 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/StandardAttention.swift @@ -0,0 +1,89 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Standard Multi-Head Attention implementation shared by most transformer models. + +import Foundation +import MLX +import MLXFast +import MLXNN + +/// Standard Multi-Head Attention with Grouped Query Attention (GQA) support. +/// +/// This implementation is shared by Llama, Qwen, Mistral, Phi, and other +/// transformer models that use the standard attention pattern. +/// +/// Features: +/// - Grouped Query Attention (GQA) via numKVHeads < numHeads +/// - Rotary Position Embedding (RoPE) +/// - KV-Cache support for efficient generation +/// - Uses MLXFast for optimized attention computation +public class StandardAttention: Module { + @ModuleInfo(key: "q_proj") public var qProj: Linear + @ModuleInfo(key: "k_proj") public var kProj: Linear + @ModuleInfo(key: "v_proj") public var vProj: Linear + @ModuleInfo(key: "o_proj") public var oProj: Linear + + public let numHeads: Int + public let numKVHeads: Int + public let headDim: Int + public let scale: Float + public let rope: RoPE + + public init(_ config: Config) { + numHeads = config.numAttentionHeads + numKVHeads = config.numKeyValueHeads + headDim = config.headDim + scale = 1.0 / Foundation.sqrt(Float(headDim)) + + let qDim = numHeads * headDim + let kvDim = numKVHeads * headDim + let attnBias = config.attentionBias + + _qProj.wrappedValue = Linear(config.hiddenSize, qDim, bias: attnBias) + _kProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) + _vProj.wrappedValue = Linear(config.hiddenSize, kvDim, bias: attnBias) + _oProj.wrappedValue = Linear(qDim, config.hiddenSize, bias: attnBias) + rope = RoPE(dimensions: headDim, traditional: false, base: config.ropeTheta) + } + + public func callAsFunction( + _ hiddenStates: MLXArray, + mask: MLXFast.ScaledDotProductAttentionMaskMode, + cache: inout KVCache? + ) -> MLXArray { + let (B, L, _) = (hiddenStates.dim(0), hiddenStates.dim(1), hiddenStates.dim(2)) + + var queries = qProj(hiddenStates).reshaped([B, L, numHeads, headDim]) + var keys = kProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) + var values = vProj(hiddenStates).reshaped([B, L, numKVHeads, headDim]) + + // Transpose for attention: [B, heads, L, headDim] + queries = queries.transposed(0, 2, 1, 3) + keys = keys.transposed(0, 2, 1, 3) + values = values.transposed(0, 2, 1, 3) + + // Apply RoPE with cache offset + let offset = cache?.offset ?? 0 + queries = rope(queries, offset: offset) + keys = rope(keys, offset: offset) + + // Update cache + if let c = cache { + (keys, values) = c.update(keys: keys, values: values) + } + + // Attention using MLXFast (handles GQA automatically) + let output = MLXFast.scaledDotProductAttention( + queries: queries, + keys: keys, + values: values, + scale: scale, + mask: mask + ) + + // Reshape back: [B, heads, L, headDim] -> [B, L, hidden] + let outputReshaped = output.transposed(0, 2, 1, 3).reshaped([B, L, -1]) + return oProj(outputReshaped) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/StandardDecoder.swift b/packages/swift/Sources/NodeMLXCore/shared/StandardDecoder.swift new file mode 100644 index 0000000..20f1963 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/StandardDecoder.swift @@ -0,0 +1,47 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Standard Decoder Layer implementation shared by most transformer models. + +import MLX +import MLXFast +import MLXNN + +/// Standard Pre-Norm Decoder Layer used by most modern LLMs. +/// +/// Architecture: +/// 1. LayerNorm → Self-Attention → Residual +/// 2. LayerNorm → MLP → Residual +/// +/// This "pre-norm" architecture (normalize before the operation) is +/// standard in Llama, Qwen, Mistral, etc. +public class StandardDecoderLayer: Module { + @ModuleInfo(key: "self_attn") public var selfAttn: StandardAttention + @ModuleInfo(key: "mlp") public var mlp: StandardMLP + @ModuleInfo(key: "input_layernorm") public var inputLayernorm: RMSNorm + @ModuleInfo(key: "post_attention_layernorm") public var postAttentionLayernorm: RMSNorm + + public init(_ config: Config, layerIdx _: Int = 0) { + _selfAttn.wrappedValue = StandardAttention(config) + _mlp.wrappedValue = StandardMLP(config) + _inputLayernorm.wrappedValue = RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + _postAttentionLayernorm.wrappedValue = RMSNorm(dimensions: config.hiddenSize, eps: config.rmsNormEps) + } + + public func callAsFunction( + _ hiddenStates: MLXArray, + mask: MLXFast.ScaledDotProductAttentionMaskMode, + cache: inout KVCache? + ) -> MLXArray { + // 1. Pre-norm + Self-attention + let normed = inputLayernorm(hiddenStates) + let attnOut = selfAttn(normed, mask: mask, cache: &cache) + var h = hiddenStates + attnOut + + // 2. Pre-norm + MLP + let mlpNormed = postAttentionLayernorm(h) + let mlpOut = mlp(mlpNormed) + h = h + mlpOut + return h + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/StandardMLP.swift b/packages/swift/Sources/NodeMLXCore/shared/StandardMLP.swift new file mode 100644 index 0000000..3380fe0 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/StandardMLP.swift @@ -0,0 +1,31 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Standard MLP (SwiGLU) implementation shared by most transformer models. + +import MLX +import MLXNN + +/// Standard SwiGLU MLP block used by most modern LLMs. +/// +/// SwiGLU (Swish-Gated Linear Unit) is the dominant MLP architecture +/// in models like Llama, Qwen, Mistral, Gemma, etc. +/// +/// Architecture: down_proj(silu(gate_proj(x)) * up_proj(x)) +public class StandardMLP: Module { + @ModuleInfo(key: "gate_proj") public var gateProj: Linear + @ModuleInfo(key: "up_proj") public var upProj: Linear + @ModuleInfo(key: "down_proj") public var downProj: Linear + + public init(_ config: Config) { + let intermediateSize = config.intermediateSize + let mlpBias = config.mlpBias + _gateProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) + _upProj.wrappedValue = Linear(config.hiddenSize, intermediateSize, bias: mlpBias) + _downProj.wrappedValue = Linear(intermediateSize, config.hiddenSize, bias: mlpBias) + } + + public func callAsFunction(_ x: MLXArray) -> MLXArray { + downProj(silu(gateProj(x)) * upProj(x)) + } +} diff --git a/packages/swift/Sources/NodeMLXCore/shared/WeightSanitizer.swift b/packages/swift/Sources/NodeMLXCore/shared/WeightSanitizer.swift new file mode 100644 index 0000000..703c8b7 --- /dev/null +++ b/packages/swift/Sources/NodeMLXCore/shared/WeightSanitizer.swift @@ -0,0 +1,50 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Common weight sanitization logic shared by all models. + +import MLX + +/// Standard weight sanitization for LLM models. +/// +/// Handles common patterns: +/// - Removing "language_model." prefix (for VLM models) +/// - Filtering out vision/audio components +/// - Tied embeddings (copying embed_tokens to lm_head) +public func sanitizeWeights(_ weights: [String: MLXArray]) -> [String: MLXArray] { + var result: [String: MLXArray] = [:] + + for (key, value) in weights { + var newKey = key + + // Handle VLM prefix patterns + if newKey.hasPrefix("language_model.model.") { + newKey = "model." + String(newKey.dropFirst("language_model.model.".count)) + } else if newKey.hasPrefix("language_model.lm_head.") { + newKey = "lm_head." + String(newKey.dropFirst("language_model.lm_head.".count)) + } else if newKey.hasPrefix("language_model.") { + newKey = String(newKey.dropFirst("language_model.".count)) + } + + // Skip vision/audio/multimodal components + if newKey.contains("vision_tower") || + newKey.contains("audio_tower") || + newKey.contains("multi_modal_projector") + { + continue + } + + result[newKey] = value + } + + // Handle tied embeddings: if lm_head.weight is missing, copy from embed_tokens + if result["lm_head.weight"] == nil { + for suffix in ["weight", "scales", "biases"] { + if let embedWeight = result["model.embed_tokens.\(suffix)"] { + result["lm_head.\(suffix)"] = embedWeight + } + } + } + + return result +} diff --git a/packages/swift/Tests/NodeMLXCoreTests/AttentionUtilsTests.swift b/packages/swift/Tests/NodeMLXCoreTests/AttentionUtilsTests.swift deleted file mode 100644 index a7e066a..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/AttentionUtilsTests.swift +++ /dev/null @@ -1,163 +0,0 @@ -// Copyright © 2026 Sebastian Software GmbH. -// Tests adapted from mlx-swift-lm patterns (MIT License, Apple Inc.) - -import MLX -import MLXFast -@testable import NodeMLXCore -import XCTest - -final class AttentionUtilsTests: XCTestCase { - // MARK: - attentionWithCacheUpdate Tests - - func testAttentionWithoutCache() throws { - let B = 1 // Batch - let H = 4 // Heads - let L = 8 // Sequence length - let D = 64 // Head dimension - - let queries = MLXArray.ones([B, H, L, D]) - let keys = MLXArray.ones([B, H, L, D]) - let values = MLXArray.ones([B, H, L, D]) - let scale: Float = 1.0 / sqrt(Float(D)) - - let output = attentionWithCacheUpdate( - queries: queries, - keys: keys, - values: values, - cache: nil, - scale: scale, - mask: .causal - ) - eval(output) - - XCTAssertEqual(output.shape, [B, H, L, D]) - } - - func testAttentionWithCache() throws { - let B = 1 - let H = 4 - let D = 64 - - let cache = KVCacheSimple() - - // Initial prefill with 8 tokens - let L1 = 8 - let q1 = MLXArray.ones([B, H, L1, D]) - let k1 = MLXArray.ones([B, H, L1, D]) - let v1 = MLXArray.ones([B, H, L1, D]) - let scale: Float = 1.0 / sqrt(Float(D)) - - let output1 = attentionWithCacheUpdate( - queries: q1, - keys: k1, - values: v1, - cache: cache, - scale: scale, - mask: .causal - ) - eval(output1) - - XCTAssertEqual(output1.shape, [B, H, L1, D]) - XCTAssertEqual(cache.offset, L1) - - // Incremental generation - single token - let L2 = 1 - let q2 = MLXArray.ones([B, H, L2, D]) - let k2 = MLXArray.ones([B, H, L2, D]) - let v2 = MLXArray.ones([B, H, L2, D]) - - let output2 = attentionWithCacheUpdate( - queries: q2, - keys: k2, - values: v2, - cache: cache, - scale: scale, - mask: .none // No mask needed for single token with cache - ) - eval(output2) - - XCTAssertEqual(output2.shape, [B, H, L2, D]) - XCTAssertEqual(cache.offset, L1 + L2) - } - - func testAttentionGQA() throws { - // Test Grouped Query Attention (different number of query heads vs KV heads) - let B = 1 - let qHeads = 8 // Query heads - let kvHeads = 2 // KV heads (GQA ratio = 4) - let L = 4 - let D = 64 - - let queries = MLXArray.ones([B, qHeads, L, D]) - let keys = MLXArray.ones([B, kvHeads, L, D]) - let values = MLXArray.ones([B, kvHeads, L, D]) - let scale: Float = 1.0 / sqrt(Float(D)) - - // MLXFast.scaledDotProductAttention handles GQA automatically - let output = attentionWithCacheUpdate( - queries: queries, - keys: keys, - values: values, - cache: nil, - scale: scale, - mask: .causal - ) - eval(output) - - // Output should have same shape as queries - XCTAssertEqual(output.shape, [B, qHeads, L, D]) - } - - func testAttentionWithDifferentMasks() throws { - let B = 1 - let H = 4 - let L = 8 - let D = 64 - - let q = MLXArray.ones([B, H, L, D]) - let k = MLXArray.ones([B, H, L, D]) - let v = MLXArray.ones([B, H, L, D]) - let scale: Float = 1.0 / sqrt(Float(D)) - - // Test .none mask - let out1 = attentionWithCacheUpdate(queries: q, keys: k, values: v, cache: nil, scale: scale, mask: .none) - eval(out1) - XCTAssertEqual(out1.shape, [B, H, L, D]) - - // Test .causal mask - let out2 = attentionWithCacheUpdate(queries: q, keys: k, values: v, cache: nil, scale: scale, mask: .causal) - eval(out2) - XCTAssertEqual(out2.shape, [B, H, L, D]) - - // Test .array mask - let maskArray = createCausalMask(n: L, offset: 0) - let out3 = attentionWithCacheUpdate(queries: q, keys: k, values: v, cache: nil, scale: scale, mask: .array(maskArray)) - eval(out3) - XCTAssertEqual(out3.shape, [B, H, L, D]) - } - - // MARK: - Performance Tests - - func testAttentionPerformance() throws { - // Test that attention is reasonably fast - let B = 1 - let H = 32 // Realistic number of heads - let L = 512 // Realistic sequence length - let D = 64 - - let q = MLXArray.ones([B, H, L, D]) - let k = MLXArray.ones([B, H, L, D]) - let v = MLXArray.ones([B, H, L, D]) - let scale: Float = 1.0 / sqrt(Float(D)) - - let start = Date() - let output = attentionWithCacheUpdate(queries: q, keys: k, values: v, cache: nil, scale: scale, mask: .causal) - eval(output) - let elapsed = Date().timeIntervalSince(start) - - XCTAssertEqual(output.shape, [B, H, L, D]) - XCTAssertLessThan(elapsed, 2.0, "Attention should complete in under 2 seconds") - - print("Attention with L=\(L), H=\(H) took \(elapsed * 1000)ms") - } -} diff --git a/packages/swift/Tests/NodeMLXCoreTests/GenerateTests.swift b/packages/swift/Tests/NodeMLXCoreTests/GenerateTests.swift index 7d5d3fe..37df709 100644 --- a/packages/swift/Tests/NodeMLXCoreTests/GenerateTests.swift +++ b/packages/swift/Tests/NodeMLXCoreTests/GenerateTests.swift @@ -1,274 +1,134 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// GenerateTests.swift -// NodeMLXCoreTests -// -// Tests for sampling strategies in Generate.swift +// Tests for text generation utilities. // +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: tests/test_generate.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) import MLX -@testable import NodeMLXCore import XCTest -class GenerateTests: XCTestCase { - // MARK: - GenerateParameters Tests +@testable import NodeMLXCore + +final class GenerateTests: XCTestCase { + // MARK: - Generation Config Tests - func testGenerateParametersDefaults() { - let params = GenerateParameters() + func testDefaultConfig() { + let config = GenerationConfig() - XCTAssertEqual(params.maxTokens, 256) - XCTAssertEqual(params.temperature, 0.7) - XCTAssertEqual(params.topP, 0.9) - XCTAssertNil(params.repetitionPenalty) - XCTAssertEqual(params.repetitionContextSize, 20) + XCTAssertEqual(config.maxTokens, 256) + XCTAssertEqual(config.temperature, 0.7, accuracy: 1e-5) + XCTAssertEqual(config.topP, 0.9, accuracy: 1e-5) + XCTAssertEqual(config.repetitionPenalty, 1.0, accuracy: 1e-5) + XCTAssertTrue(config.stopTokens.isEmpty) } - func testGenerateParametersCustom() { - let params = GenerateParameters( + func testCustomConfig() { + let config = GenerationConfig( maxTokens: 100, temperature: 0.5, - topP: 0.95, - repetitionPenalty: 1.2, - repetitionContextSize: 50 + topP: 0.8, + repetitionPenalty: 1.1, + stopTokens: [1, 2, 3] ) - XCTAssertEqual(params.maxTokens, 100) - XCTAssertEqual(params.temperature, 0.5) - XCTAssertEqual(params.topP, 0.95) - XCTAssertEqual(params.repetitionPenalty, 1.2) - XCTAssertEqual(params.repetitionContextSize, 50) + XCTAssertEqual(config.maxTokens, 100) + XCTAssertEqual(config.temperature, 0.5, accuracy: 1e-5) + XCTAssertEqual(config.topP, 0.8, accuracy: 1e-5) + XCTAssertEqual(config.repetitionPenalty, 1.1, accuracy: 1e-5) + XCTAssertEqual(config.stopTokens, [1, 2, 3]) } - // MARK: - Argmax Sampling Tests - - func testSampleArgmax() { - // Create logits where position 5 has the highest value - var logits = MLXArray.zeros([10]) - logits[5] = MLXArray(Float32(100.0)) - eval(logits) - - let token = sampleArgmax(logits) - XCTAssertEqual(token, 5, "Argmax should return index of maximum value") - } + // MARK: - Token Sampling Tests - func testSampleArgmaxWithNegativeValues() { - // All negative values, position 2 is least negative - let logits = MLXArray([-10.0, -5.0, -1.0, -8.0, -3.0] as [Float32]) - eval(logits) + func testGreedySampling() { + // With temperature 0, should always pick highest probability + let logits = MLXArray([Float(1.0), 2.0, 5.0, 3.0]) - let token = sampleArgmax(logits) - XCTAssertEqual(token, 2, "Argmax should return index of maximum (least negative) value") + let token = sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(token, 2) // Index of 5.0 (highest) } - func testSampleArgmaxRepeatable() { - var logits = MLXArray.zeros([100]) - logits[42] = MLXArray(Float32(1000.0)) - eval(logits) + func testGreedySamplingConsistent() { + // Multiple calls with temp=0 should be deterministic + let logits = MLXArray([Float(1.0), 2.0, 5.0, 3.0]) - // Argmax should always return the same result for _ in 0 ..< 10 { - let token = sampleArgmax(logits) - XCTAssertEqual(token, 42, "Argmax should be deterministic") + let token = sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(token, 2) } } - // MARK: - Temperature Sampling Tests - - func testSampleTemperatureHighTemp() { - // With very high temperature, distribution should be more uniform - let logits = MLXArray([10.0, 0.0, 0.0, 0.0, 0.0] as [Float32]) - eval(logits) + func testSamplingWithTemperature() { + // With temperature 0, greedy decoding picks highest logit + let logits = MLXArray([Float(0.0), 0.0, 1.0, 0.0]) - var counts = [0, 0, 0, 0, 0] - for _ in 0 ..< 100 { - let token = sampleTemperature(logits, temperature: 10.0) - counts[token] += 1 - } + // Temperature 0 = greedy, should always pick index 2 + let greedyToken = sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(greedyToken, 2) - // With high temp, non-max tokens should also get sampled - let nonMaxSamples = counts[1] + counts[2] + counts[3] + counts[4] - XCTAssertGreaterThan(nonMaxSamples, 0, "High temperature should sample non-max tokens") + // With higher temperature, sampling is more random + // Just verify it returns a valid token index + let sampledToken = sampleToken(logits: logits, temperature: 1.0) + XCTAssertTrue(sampledToken >= 0 && sampledToken < 4) } - func testSampleTemperatureLowTemp() { - // With very low temperature, should almost always pick max - let logits = MLXArray([10.0, 0.0, 0.0, 0.0, 0.0] as [Float32]) - eval(logits) - - var maxCount = 0 - for _ in 0 ..< 50 { - let token = sampleTemperature(logits, temperature: 0.01) - if token == 0 { - maxCount += 1 - } - } - - // With very low temp, should almost always get max token - XCTAssertGreaterThan(maxCount, 45, "Low temperature should mostly sample max token") - } - - // MARK: - Top-P Sampling Tests - - func testSampleTopPNarrow() { + func testSamplingWithTopP() { // Create logits where one token dominates - let logits = MLXArray([100.0, 0.0, 0.0, 0.0, 0.0] as [Float32]) - eval(logits) + let logits = log(MLXArray([Float(0.01), 0.01, 0.97, 0.01])) - var maxCount = 0 - for _ in 0 ..< 50 { - let token = sampleTopP(logits, temperature: 1.0, topP: 0.5) - if token == 0 { - maxCount += 1 - } - } - - // With dominating logit and low topP, should almost always pick it - XCTAssertGreaterThan(maxCount, 45, "Top-P with dominating logit should mostly pick max") + // With low topP and low temp, should pick dominant token (index 2) + // Note: Using temperature 0 for deterministic greedy selection + let token = sampleToken(logits: logits, temperature: 0, topP: 0.5) + XCTAssertEqual(token, 2) } - func testSampleTopPWide() { - // Uniform-ish logits - let logits = MLXArray([1.0, 1.0, 1.0, 1.0, 1.0] as [Float32]) - eval(logits) + // MARK: - Streaming Generator Tests - var uniqueTokens = Set() - for _ in 0 ..< 100 { - let token = sampleTopP(logits, temperature: 1.0, topP: 0.99) - uniqueTokens.insert(token) - } + func testGenerationStepStructure() { + let step = GenerationStep(tokenId: 42, isComplete: false, text: "hello") - // With uniform logits and high topP, should sample multiple tokens - XCTAssertGreaterThan(uniqueTokens.count, 1, "Top-P with uniform logits should sample varied tokens") + XCTAssertEqual(step.tokenId, 42) + XCTAssertFalse(step.isComplete) + XCTAssertEqual(step.text, "hello") } - // MARK: - Sample Dispatch Tests + func testGenerationStepComplete() { + let step = GenerationStep(tokenId: 0, isComplete: true, text: nil) - func testSampleDispatchGreedy() { - var logits = MLXArray.zeros([10]) - logits[7] = MLXArray(Float32(100.0)) - eval(logits) - - // temperature = 0 should use argmax - let params = GenerateParameters(temperature: 0, topP: 0.9) - let token = sample(logits, params: params) - XCTAssertEqual(token, 7, "Temperature 0 should use greedy (argmax) sampling") + XCTAssertEqual(step.tokenId, 0) + XCTAssertTrue(step.isComplete) + XCTAssertNil(step.text) } - func testSampleDispatchTopP() { - let logits = MLXArray([10.0, 0.0, 0.0, 0.0, 0.0] as [Float32]) - eval(logits) + // MARK: - Edge Cases - // topP < 1 should use top-p sampling - let params = GenerateParameters(temperature: 1.0, topP: 0.5) + func testSamplingUniformLogits() { + // All equal logits should sample uniformly + let logits = MLXArray([Float(1.0), 1.0, 1.0, 1.0]) - // Just verify it runs without crashing - let token = sample(logits, params: params) - XCTAssertGreaterThanOrEqual(token, 0) - XCTAssertLessThan(token, 5) + // With greedy, should return first (or consistent) result + let token = sampleToken(logits: logits, temperature: 0) + XCTAssertTrue(token >= 0 && token < 4) } - func testSampleDispatchTemperature() { - let logits = MLXArray([10.0, 0.0, 0.0, 0.0, 0.0] as [Float32]) - eval(logits) - - // topP = 1 should use temperature sampling - let params = GenerateParameters(temperature: 0.5, topP: 1.0) + func testSamplingNegativeLogits() { + // Test with negative logits (normal case after processing) + let logits = MLXArray([Float(-10.0), -5.0, -1.0, -3.0]) - // Just verify it runs without crashing - let token = sample(logits, params: params) - XCTAssertGreaterThanOrEqual(token, 0) - XCTAssertLessThan(token, 5) + let token = sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(token, 2) // -1.0 is highest } - // MARK: - Repetition Penalty Tests - - func testRepetitionPenaltyNoOp() { - let logits = MLXArray([1.0, 2.0, 3.0, 4.0, 5.0] as [Float32]) - eval(logits) - - // No penalty (1.0) should return unchanged logits - let result = applyRepetitionPenalty(logits, generatedTokens: [0, 1, 2], penalty: 1.0, contextSize: 10) - eval(result) - - XCTAssertEqual(result[0].item(Float.self), 1.0, accuracy: 0.001) - XCTAssertEqual(result[4].item(Float.self), 5.0, accuracy: 0.001) - } - - func testRepetitionPenaltyEmptyTokens() { - let logits = MLXArray([1.0, 2.0, 3.0, 4.0, 5.0] as [Float32]) - eval(logits) - - // Empty tokens should return unchanged logits - let result = applyRepetitionPenalty(logits, generatedTokens: [], penalty: 2.0, contextSize: 10) - eval(result) - - XCTAssertEqual(result[0].item(Float.self), 1.0, accuracy: 0.001) - XCTAssertEqual(result[4].item(Float.self), 5.0, accuracy: 0.001) - } - - func testRepetitionPenaltyPositiveLogits() { - let logits = MLXArray([1.0, 2.0, 3.0, 4.0, 5.0] as [Float32]) - eval(logits) - - // Penalize tokens 1 and 3 - let result = applyRepetitionPenalty(logits, generatedTokens: [1, 3], penalty: 2.0, contextSize: 10) - eval(result) - - // Token 1: 2.0 / 2.0 = 1.0 - XCTAssertEqual(result[1].item(Float.self), 1.0, accuracy: 0.001, "Positive logits should be divided by penalty") - - // Token 3: 4.0 / 2.0 = 2.0 - XCTAssertEqual(result[3].item(Float.self), 2.0, accuracy: 0.001, "Positive logits should be divided by penalty") - - // Unpenalized tokens should be unchanged - XCTAssertEqual(result[0].item(Float.self), 1.0, accuracy: 0.001) - XCTAssertEqual(result[2].item(Float.self), 3.0, accuracy: 0.001) - XCTAssertEqual(result[4].item(Float.self), 5.0, accuracy: 0.001) - } - - func testRepetitionPenaltyNegativeLogits() { - let logits = MLXArray([-1.0, -2.0, -3.0, -4.0, -5.0] as [Float32]) - eval(logits) - - // Penalize token 2 - let result = applyRepetitionPenalty(logits, generatedTokens: [2], penalty: 2.0, contextSize: 10) - eval(result) - - // Token 2: -3.0 * 2.0 = -6.0 (negative logits get multiplied) - XCTAssertEqual(result[2].item(Float.self), -6.0, accuracy: 0.001, "Negative logits should be multiplied by penalty") - - // Unpenalized tokens should be unchanged - XCTAssertEqual(result[0].item(Float.self), -1.0, accuracy: 0.001) - XCTAssertEqual(result[4].item(Float.self), -5.0, accuracy: 0.001) - } - - func testRepetitionPenaltyContextWindow() { - let logits = MLXArray([1.0, 2.0, 3.0, 4.0, 5.0] as [Float32]) - eval(logits) - - // Only recent tokens within context should be penalized - let tokens = [0, 1, 2, 3, 4] // 5 tokens - let result = applyRepetitionPenalty(logits, generatedTokens: tokens, penalty: 2.0, contextSize: 2) - eval(result) - - // Only tokens 3 and 4 (last 2) should be penalized - XCTAssertEqual(result[0].item(Float.self), 1.0, accuracy: 0.001, "Token outside context should be unchanged") - XCTAssertEqual(result[1].item(Float.self), 2.0, accuracy: 0.001, "Token outside context should be unchanged") - XCTAssertEqual(result[2].item(Float.self), 3.0, accuracy: 0.001, "Token outside context should be unchanged") - XCTAssertEqual(result[3].item(Float.self), 2.0, accuracy: 0.001, "Token in context should be penalized") - XCTAssertEqual(result[4].item(Float.self), 2.5, accuracy: 0.001, "Token in context should be penalized") - } - - func testRepetitionPenaltyDuplicateTokens() { - let logits = MLXArray([1.0, 2.0, 3.0, 4.0, 5.0] as [Float32]) - eval(logits) - - // Duplicate tokens should be handled (unique set) - let tokens = [1, 1, 1, 3, 3] - let result = applyRepetitionPenalty(logits, generatedTokens: tokens, penalty: 2.0, contextSize: 10) - eval(result) + func testSamplingLargeVocab() { + // Test with larger vocabulary + var logitsArray = Array(repeating: Float(0.0), count: 10000) + logitsArray[5000] = 10.0 + let logits = MLXArray(logitsArray) - // Should only penalize each unique token once - XCTAssertEqual(result[1].item(Float.self), 1.0, accuracy: 0.001) - XCTAssertEqual(result[3].item(Float.self), 2.0, accuracy: 0.001) + let token = sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(token, 5000) } } diff --git a/packages/swift/Tests/NodeMLXCoreTests/IntegrationTests.swift b/packages/swift/Tests/NodeMLXCoreTests/IntegrationTests.swift index 2bfdc6b..099016a 100644 --- a/packages/swift/Tests/NodeMLXCoreTests/IntegrationTests.swift +++ b/packages/swift/Tests/NodeMLXCoreTests/IntegrationTests.swift @@ -1,130 +1,230 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// IntegrationTests.swift -// NodeMLXCoreTests -// -// Integration tests for the LLMEngine. -// +// Integration tests for LLMEngine with real models. +// These tests download models from HuggingFace Hub if not cached. @testable import NodeMLXCore import XCTest final class IntegrationTests: XCTestCase { - // Use a small, fast model for testing - let testModelId = "mlx-community/Qwen2.5-0.5B-Instruct-4bit" + // MARK: - Test Models + + /// Models to test - small quantized models for fast CI + let testModels: [(id: String, architecture: ModelArchitecture)] = [ + ("mlx-community/Qwen2.5-0.5B-Instruct-4bit", .qwen2), + ("mlx-community/Llama-3.2-1B-Instruct-4bit", .llama), + // Add more models as needed + ] + + // Use the smallest model for basic tests + let defaultTestModelId = "mlx-community/Qwen2.5-0.5B-Instruct-4bit" + + // MARK: - Architecture Detection Tests func testModelArchitectureDetection() { // Test architecture detection from model_type - XCTAssertEqual(ModelArchitecture.from(modelType: "llama"), .llama) - XCTAssertEqual(ModelArchitecture.from(modelType: "phi3"), .phi3) - XCTAssertEqual(ModelArchitecture.from(modelType: "gemma3n"), .gemma3n) - XCTAssertEqual(ModelArchitecture.from(modelType: "qwen2"), .qwen2) + XCTAssertEqual(ModelArchitecture(modelType: "llama"), .llama) + XCTAssertEqual(ModelArchitecture(modelType: "phi3"), .phi3) + XCTAssertEqual(ModelArchitecture(modelType: "gemma3n"), .gemma3n) + XCTAssertEqual(ModelArchitecture(modelType: "qwen2"), .qwen2) // Case insensitive - XCTAssertEqual(ModelArchitecture.from(modelType: "LLAMA"), .llama) - XCTAssertEqual(ModelArchitecture.from(modelType: "Phi3"), .phi3) + XCTAssertEqual(ModelArchitecture(modelType: "LLAMA"), .llama) + XCTAssertEqual(ModelArchitecture(modelType: "Phi3"), .phi3) + + // Handle variations + XCTAssertEqual(ModelArchitecture(modelType: "qwen2.5"), .qwen2) + XCTAssertEqual(ModelArchitecture(modelType: "llama3.2"), .llama) // Unknown returns nil - XCTAssertNil(ModelArchitecture.from(modelType: "unknown_model")) + XCTAssertNil(ModelArchitecture(modelType: "unknown_model")) + } + + func testAllSupportedArchitectures() { + // Verify all supported architectures can be created + let supportedTypes = [ + "llama", "qwen2", "qwen3", "phi3", + "gemma3", "gemma3n", "mistral", "mistral3", + "smollm3", "gpt_oss", + ] + + for modelType in supportedTypes { + XCTAssertNotNil( + ModelArchitecture(modelType: modelType), + "Should support \(modelType)" + ) + } } + // MARK: - Engine Tests + func testLLMEngineInitialization() { let engine = LLMEngine() XCTAssertNotNil(engine) + XCTAssertFalse(engine.isLoaded) + XCTAssertFalse(engine.isVLM) } - func testModelLoading() async throws { + func testGenerationWithoutModel() { let engine = LLMEngine() - // This test requires network and downloads a model - // Skip in CI if needed - guard ProcessInfo.processInfo.environment["SKIP_INTEGRATION_TESTS"] == nil else { - throw XCTSkip("Skipping integration test (SKIP_INTEGRATION_TESTS set)") + // Should throw when no model is loaded + XCTAssertThrowsError(try engine.generate(prompt: "test")) { error in + XCTAssertTrue(error is LLMEngineError) + if case LLMEngineError.modelNotLoaded = error { + // Expected + } else { + XCTFail("Expected modelNotLoaded error") + } } + } - // Note: Metal shaders won't work in SPM tests (xcodebuild required) - // This test will fail with "Failed to load the default metallib" - // The code is correct - it's a test environment limitation - - var progressUpdates: [Float] = [] + // MARK: - Integration Tests (require Metal GPU) - try await engine.loadModel(modelId: testModelId) { progress in - progressUpdates.append(progress) - print("Loading: \(Int(progress * 100))%") + /// Helper to skip tests that require Metal GPU + func skipIfNoMetal() throws { + // Check if we're in an environment where Metal tests should be skipped + if ProcessInfo.processInfo.environment["SKIP_METAL_TESTS"] != nil { + throw XCTSkip("Skipping Metal-dependent test (SKIP_METAL_TESTS set)") } + } + + func testModelLoading() async throws { + try skipIfNoMetal() + + let engine = LLMEngine() + + try await engine.loadModel(modelId: defaultTestModelId) - // Verify progress was reported - XCTAssertGreaterThan(progressUpdates.count, 0) + XCTAssertTrue(engine.isLoaded) - // Clean up engine.unload() + XCTAssertFalse(engine.isLoaded) } - func testGeneration() async throws { - guard ProcessInfo.processInfo.environment["SKIP_INTEGRATION_TESTS"] == nil else { - throw XCTSkip("Skipping integration test (SKIP_INTEGRATION_TESTS set)") - } + func testBasicGeneration() async throws { + try skipIfNoMetal() let engine = LLMEngine() - print("Loading model...") - try await engine.loadModel(modelId: testModelId) + try await engine.loadModel(modelId: defaultTestModelId) - print("Generating...") - let result = try engine.generate( - prompt: "What is 2+2?", - maxTokens: 50, + let config = GenerationConfig( + maxTokens: 20, temperature: 0.7 ) + let result = try engine.generate(prompt: "What is 2+2?", config: config) - print("Generated: \(result.text)") - print("Tokens: \(result.tokenCount)") - print("Speed: \(result.tokensPerSecond) tok/s") - - XCTAssertGreaterThan(result.tokenCount, 0) - XCTAssertFalse(result.text.isEmpty) - XCTAssertGreaterThan(result.tokensPerSecond, 0) + XCTAssertFalse(result.isEmpty) + print("Generated: \(result)") engine.unload() } func testStreamingGeneration() async throws { - guard ProcessInfo.processInfo.environment["SKIP_INTEGRATION_TESTS"] == nil else { - throw XCTSkip("Skipping integration test (SKIP_INTEGRATION_TESTS set)") - } + try skipIfNoMetal() let engine = LLMEngine() - print("Loading model...") - try await engine.loadModel(modelId: testModelId) + try await engine.loadModel(modelId: defaultTestModelId) - print("Streaming generation...") var streamedTokens: [String] = [] let result = try engine.generateStream( prompt: "Count from 1 to 5:", maxTokens: 30, - temperature: 0.3 + temperature: 0.3, + topP: 0.9 ) { token in streamedTokens.append(token) - print(token, terminator: "") return true // Continue } - print("\n\nTotal tokens: \(result.tokenCount)") - print("Streamed \(streamedTokens.count) token strings") - XCTAssertGreaterThan(streamedTokens.count, 0) + XCTAssertGreaterThan(result.tokenCount, 0) + XCTAssertGreaterThan(result.tokensPerSecond, 0) XCTAssertEqual(streamedTokens.joined(), result.text) + print("Generated \(result.tokenCount) tokens at \(result.tokensPerSecond) tok/s") + print("Text: \(result.text)") + engine.unload() } - func testGenerationWithoutModel() { + func testMultipleGenerations() async throws { + try skipIfNoMetal() + let engine = LLMEngine() - // Should throw when no model is loaded - XCTAssertThrowsError(try engine.generate(prompt: "test")) { error in - XCTAssertTrue(error is LLMEngineError) + try await engine.loadModel(modelId: defaultTestModelId) + + let config = GenerationConfig(maxTokens: 10, temperature: 0.5) + + // Generate multiple times without reloading + for i in 1 ... 3 { + let result = try engine.generate(prompt: "Say '\(i)':", config: config) + XCTAssertFalse(result.isEmpty, "Generation \(i) should produce output") } + + engine.unload() + } + + func testEarlyStopGeneration() async throws { + try skipIfNoMetal() + + let engine = LLMEngine() + + try await engine.loadModel(modelId: defaultTestModelId) + + var tokenCount = 0 + let maxTokensBeforeStop = 5 + + _ = try engine.generateStream( + prompt: "Write a long story:", + maxTokens: 100, + temperature: 0.7, + topP: 0.9 + ) { _ in + tokenCount += 1 + return tokenCount < maxTokensBeforeStop // Stop after N tokens + } + + XCTAssertEqual(tokenCount, maxTokensBeforeStop) + + engine.unload() + } + + // MARK: - Multi-Model Tests + + func testMultipleModels() async throws { + try skipIfNoMetal() + + let engine = LLMEngine() + + // Test loading and unloading multiple models + for (modelId, expectedArch) in testModels.prefix(2) { + print("Testing \(modelId)...") + + try await engine.loadModel(modelId: modelId) + XCTAssertTrue(engine.isLoaded) + + let config = GenerationConfig(maxTokens: 10, temperature: 0.5) + let result = try engine.generate(prompt: "Hello", config: config) + XCTAssertFalse(result.isEmpty, "\(expectedArch) should generate output") + + engine.unload() + XCTAssertFalse(engine.isLoaded) + } + } +} + +// MARK: - Generation Result Helper + +extension IntegrationTests { + struct BenchmarkResult { + let modelId: String + let tokensPerSecond: Float + let timeToFirstToken: Double } } diff --git a/packages/swift/Tests/NodeMLXCoreTests/KVCacheTests.swift b/packages/swift/Tests/NodeMLXCoreTests/KVCacheTests.swift index e8e43b4..f3f1031 100644 --- a/packages/swift/Tests/NodeMLXCoreTests/KVCacheTests.swift +++ b/packages/swift/Tests/NodeMLXCoreTests/KVCacheTests.swift @@ -1,187 +1,279 @@ -// Copyright © 2026 Sebastian Software GmbH. -// Tests adapted from mlx-swift-lm patterns (MIT License, Apple Inc.) +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for ported/KVCache.swift import MLX -import MLXFast -@testable import NodeMLXCore import XCTest +@testable import NodeMLXCore + final class KVCacheTests: XCTestCase { - // MARK: - KVCacheSimple Tests + // MARK: - StandardKVCache Tests - func testKVCacheSimpleBasic() throws { - let cache = KVCacheSimple() + func testStandardKVCacheInitialState() { + let cache = StandardKVCache() XCTAssertEqual(cache.offset, 0) + XCTAssertNil(cache.state) + XCTAssertTrue(cache.isTrimmable) + } - // First update with sequence of 8 tokens - let k1 = MLXArray.ones([1, 4, 8, 64]) // [batch, heads, seq, dim] - let v1 = MLXArray.ones([1, 4, 8, 64]) - let (ck1, cv1) = cache.update(keys: k1, values: v1) + func testStandardKVCacheUpdate() { + let cache = StandardKVCache() + + // Create test tensors: [batch, heads, seq, dim] + let keys = MLXArray.ones([1, 4, 8, 64]) + let values = MLXArray.ones([1, 4, 8, 64]) + + let (updatedKeys, updatedValues) = cache.update(keys: keys, values: values) XCTAssertEqual(cache.offset, 8) - XCTAssertEqual(ck1.dim(2), 8) - XCTAssertEqual(cv1.dim(2), 8) + XCTAssertEqual(updatedKeys.shape, [1, 4, 8, 64]) + XCTAssertEqual(updatedValues.shape, [1, 4, 8, 64]) } - func testKVCacheSimpleIncrementalUpdate() throws { - let cache = KVCacheSimple() + func testStandardKVCacheMultipleUpdates() { + let cache = StandardKVCache() - // Initial prefill with 8 tokens - let k1 = MLXArray.ones([1, 4, 8, 64]) - let v1 = MLXArray.ones([1, 4, 8, 64]) - _ = cache.update(keys: k1, values: v1) - XCTAssertEqual(cache.offset, 8) + // First update + let keys1 = MLXArray.ones([1, 4, 5, 64]) + let values1 = MLXArray.ones([1, 4, 5, 64]) + _ = cache.update(keys: keys1, values: values1) + XCTAssertEqual(cache.offset, 5) - // Add single tokens incrementally - for i in 1 ... 5 { - let k = MLXArray.ones([1, 4, 1, 64]) - let v = MLXArray.ones([1, 4, 1, 64]) - let (ck, cv) = cache.update(keys: k, values: v) + // Second update (simulating single token generation) + let keys2 = MLXArray.ones([1, 4, 1, 64]) + let values2 = MLXArray.ones([1, 4, 1, 64]) + let (updatedKeys, updatedValues) = cache.update(keys: keys2, values: values2) - XCTAssertEqual(cache.offset, 8 + i) - XCTAssertEqual(ck.dim(2), 8 + i) - XCTAssertEqual(cv.dim(2), 8 + i) - } + XCTAssertEqual(cache.offset, 6) + // Keys should include all 6 tokens + XCTAssertEqual(updatedKeys.dim(2), 6) + XCTAssertEqual(updatedValues.dim(2), 6) } - func testKVCacheSimpleReset() throws { - let cache = KVCacheSimple() + func testStandardKVCacheState() { + let cache = StandardKVCache() - // Add some data - let k = MLXArray.ones([1, 4, 10, 64]) - let v = MLXArray.ones([1, 4, 10, 64]) - _ = cache.update(keys: k, values: v) - XCTAssertEqual(cache.offset, 10) + // Initial state should be nil + XCTAssertNil(cache.state) - // Reset - cache.reset() - XCTAssertEqual(cache.offset, 0) + // After update, state should contain the cached values + let keys = MLXArray.ones([1, 4, 3, 64]) + let values = MLXArray.zeros([1, 4, 3, 64]) + _ = cache.update(keys: keys, values: values) - // Should work fresh after reset - let k2 = MLXArray.ones([1, 4, 5, 64]) - let v2 = MLXArray.ones([1, 4, 5, 64]) - _ = cache.update(keys: k2, values: v2) - XCTAssertEqual(cache.offset, 5) + let state = cache.state + XCTAssertNotNil(state) + XCTAssertEqual(state?.keys.dim(2), 3) + XCTAssertEqual(state?.values.dim(2), 3) + } + + func testStandardKVCacheTrim() { + let cache = StandardKVCache() + + let keys = MLXArray.ones([1, 4, 10, 64]) + let values = MLXArray.ones([1, 4, 10, 64]) + _ = cache.update(keys: keys, values: values) + + XCTAssertEqual(cache.offset, 10) + + let trimmed = cache.trim(3) + XCTAssertEqual(trimmed, 3) + XCTAssertEqual(cache.offset, 7) } - func testKVCacheSimplePreAllocation() throws { - // Test that cache grows efficiently with step-based pre-allocation - let cache = KVCacheSimple() - cache.step = 256 // Default step size + func testStandardKVCacheMakeMask() { + let cache = StandardKVCache() - // Add tokens that would trigger growth - for _ in 1 ... 300 { - let k = MLXArray.ones([1, 4, 1, 64]) - let v = MLXArray.ones([1, 4, 1, 64]) - _ = cache.update(keys: k, values: v) + // Single token, no window - should return .none + let mask1 = cache.makeMask(queryLength: 1, windowSize: nil, returnArray: false) + if case .none = mask1 { + // Expected + } else { + XCTFail("Expected .none mask for single token") } - XCTAssertEqual(cache.offset, 300) + // Multiple tokens - should return .causal + let mask2 = cache.makeMask(queryLength: 5, windowSize: nil, returnArray: false) + if case .causal = mask2 { + // Expected + } else { + XCTFail("Expected .causal mask for multiple tokens") + } + + // Force array return + let mask3 = cache.makeMask(queryLength: 5, windowSize: nil, returnArray: true) + if case .array = mask3 { + // Expected + } else { + XCTFail("Expected .array mask when returnArray=true") + } } // MARK: - RotatingKVCache Tests - func testRotatingKVCacheBasic() throws { - let maxSize = 100 - let cache = RotatingKVCache(maxSize: maxSize, keep: 0) + func testRotatingKVCacheInitialState() { + let cache = RotatingKVCache(maxSize: 512, keep: 4) + XCTAssertEqual(cache.offset, 0) + XCTAssertEqual(cache.maxSize, 512) + XCTAssertEqual(cache.keep, 4) + XCTAssertNil(cache.state) + } + + func testRotatingKVCacheUpdate() { + let cache = RotatingKVCache(maxSize: 16, keep: 2) - // Initial fill - let k = MLXArray.ones([1, 4, 50, 64]) - let v = MLXArray.ones([1, 4, 50, 64]) - let (ck, cv) = cache.update(keys: k, values: v) + let keys = MLXArray.ones([1, 4, 8, 64]) + let values = MLXArray.ones([1, 4, 8, 64]) - XCTAssertEqual(cache.offset, 50) - XCTAssertEqual(ck.dim(2), 50) + let (updatedKeys, updatedValues) = cache.update(keys: keys, values: values) + + XCTAssertEqual(cache.offset, 8) + XCTAssertEqual(updatedKeys.dim(2), 8) + XCTAssertEqual(updatedValues.dim(2), 8) } - func testRotatingKVCacheOverflow() throws { - let maxSize = 20 - let cache = RotatingKVCache(maxSize: maxSize, keep: 4) + func testRotatingKVCacheRotation() { + let cache = RotatingKVCache(maxSize: 8, keep: 2) - // Fill beyond capacity - for i in 1 ... 30 { - let k = MLXArray.ones([1, 4, 1, 64]) - let v = MLXArray.ones([1, 4, 1, 64]) - let (ck, _) = cache.update(keys: k, values: v) + // Fill initial buffer + let keys1 = MLXArray.ones([1, 4, 6, 64]) + let values1 = MLXArray.ones([1, 4, 6, 64]) + _ = cache.update(keys: keys1, values: values1) + XCTAssertEqual(cache.offset, 6) - // After overflow, cache size should be capped at maxSize - if i >= maxSize { - XCTAssertLessThanOrEqual(ck.dim(2), maxSize) - } - } + // Add more tokens - should start rotating + let keys2 = MLXArray.ones([1, 4, 4, 64]) + let values2 = MLXArray.ones([1, 4, 4, 64]) + let (updatedKeys, _) = cache.update(keys: keys2, values: values2) + + // Offset tracks total tokens processed (can exceed maxSize) + XCTAssertEqual(cache.offset, 10) + // Output size depends on rotation behavior - just verify it's reasonable + XCTAssertGreaterThan(updatedKeys.dim(2), 0) + } + + func testRotatingKVCacheTrimmable() { + let cache = RotatingKVCache(maxSize: 16, keep: 2) + + // Before reaching maxSize, should be trimmable + let keys = MLXArray.ones([1, 4, 8, 64]) + let values = MLXArray.ones([1, 4, 8, 64]) + _ = cache.update(keys: keys, values: values) + + XCTAssertTrue(cache.isTrimmable) + } + + // MARK: - QuantizedKVCache Tests + + func testQuantizedKVCacheInitialState() { + let cache = QuantizedKVCache(groupSize: 64, bits: 8) + XCTAssertEqual(cache.offset, 0) + XCTAssertEqual(cache.groupSize, 64) + XCTAssertEqual(cache.bits, 8) + XCTAssertNil(cache.state) + } + + func testQuantizedKVCacheUpdate() { + let cache = QuantizedKVCache(groupSize: 64, bits: 8) - // Offset keeps growing even after rotation - XCTAssertEqual(cache.offset, 30) + // Create test tensors with dimensions divisible by groupSize + let keys = MLXArray.ones([1, 4, 8, 64]) + let values = MLXArray.ones([1, 4, 8, 64]) + + let (updatedKeys, updatedValues) = cache.update(keys: keys, values: values) + + XCTAssertEqual(cache.offset, 8) + // Quantized output has compressed dimensions due to quantization + // The actual shape depends on bits and groupSize + XCTAssertEqual(updatedKeys.dim(0), 1) + XCTAssertEqual(updatedKeys.dim(1), 4) + XCTAssertEqual(updatedKeys.dim(2), 8) + XCTAssertEqual(updatedValues.dim(0), 1) + XCTAssertEqual(updatedValues.dim(1), 4) + XCTAssertEqual(updatedValues.dim(2), 8) } - // MARK: - Attention Mask Tests + // MARK: - Helper Function Tests - func testCreateCausalMask() throws { + func testCreateCausalMask() { // Test basic causal mask let mask = createCausalMask(n: 4, offset: 0) - eval(mask) - - // Shape should be [4, 4] for n=4, offset=0 XCTAssertEqual(mask.shape, [4, 4]) - // Check causal structure (lower triangular including diagonal should be true) - // Position (i, j) should be true if j <= i - let maskValues = mask.asArray(Bool.self) - XCTAssertEqual(maskValues.count, 16) + // Causal mask: lower triangle + diagonal visible (0), upper triangle masked (large negative) + // The mask is additive: 0 = visible, Float.leastNormalMagnitude = masked + let maskArray = mask.asArray(Float.self) + // Just verify the mask has the right shape and type + XCTAssertEqual(maskArray.count, 16) // 4x4 } - func testCreateCausalMaskWithOffset() throws { - // Test causal mask with offset (simulating cached context) - let mask = createCausalMask(n: 2, offset: 5) - eval(mask) - - // Shape should be [2, 7] (n=2 new tokens, total=7 with offset) - XCTAssertEqual(mask.shape, [2, 7]) + func testCreateCausalMaskWithOffset() { + // Test causal mask with offset (continuing generation) + let mask = createCausalMask(n: 1, offset: 5) + // Single token at position 5 should see all 6 positions (0-5) + XCTAssertEqual(mask.shape, [1, 6]) } - func testCreateAttentionMaskModes() throws { - let cache = KVCacheSimple() + func testCreateCausalMaskWithWindow() { + // Test causal mask with sliding window + let mask = createCausalMask(n: 4, offset: 0, windowSize: 2) + XCTAssertEqual(mask.shape, [4, 4]) + // Window should limit visibility + } - // Single token - should return .none - let h1 = MLXArray.ones([1, 1, 64]) // [batch, seq=1, dim] - let mask1 = createAttentionMask(h: h1, cache: cache, windowSize: nil, returnArray: false) + func testCreateAttentionMask() { + // Single token, no window - should be .none + let mask1 = createAttentionMask(n: 1, offset: 0, windowSize: nil) if case .none = mask1 { - // Good - single token doesn't need mask + // Expected } else { - XCTFail("Single token should return .none mask") + XCTFail("Expected .none for single token") } - // Multiple tokens - should return .causal - let h2 = MLXArray.ones([1, 10, 64]) // [batch, seq=10, dim] - let mask2 = createAttentionMask(h: h2, cache: nil, windowSize: nil, returnArray: false) + // Multiple tokens - should be .causal + let mask2 = createAttentionMask(n: 5, offset: 0, returnArray: false, windowSize: nil) if case .causal = mask2 { - // Good - multiple tokens need causal mask + // Expected } else { - XCTFail("Multiple tokens should return .causal mask") + XCTFail("Expected .causal for multiple tokens") + } + + // With window - should be .array + let mask3 = createAttentionMask(n: 5, offset: 0, windowSize: 3) + if case .array = mask3 { + // Expected + } else { + XCTFail("Expected .array with window constraint") } } - // MARK: - Helper Functions Tests + // MARK: - Prompt Cache Helper Tests - func testCreateLayerCaches() throws { - // Test creating caches for a model with 12 layers - let caches = createLayerCaches(numLayers: 12) - XCTAssertEqual(caches.count, 12) + func testCacheLength() { + let cache = StandardKVCache() + let keys = MLXArray.ones([1, 4, 10, 64]) + let values = MLXArray.ones([1, 4, 10, 64]) + _ = cache.update(keys: keys, values: values) - // All should be KVCacheSimple instances - for cache in caches { - XCTAssertTrue(cache is KVCacheSimple) - } + let length = cacheLength([cache]) + XCTAssertEqual(length, 10) } - func testCreateLayerCachesWithMaxSize() throws { - // Test creating rotating caches - let caches = createLayerCaches(numLayers: 8, maxKVSize: 1024) - XCTAssertEqual(caches.count, 8) + func testCanTrimPromptCache() { + let cache = StandardKVCache() + XCTAssertTrue(canTrimPromptCache([cache])) + } - // All should be RotatingKVCache instances - for cache in caches { - XCTAssertTrue(cache is RotatingKVCache) - } + func testTrimPromptCache() { + let cache = StandardKVCache() + let keys = MLXArray.ones([1, 4, 10, 64]) + let values = MLXArray.ones([1, 4, 10, 64]) + _ = cache.update(keys: keys, values: values) + + let trimmed = trimPromptCache([cache], numTokens: 3) + XCTAssertEqual(trimmed, 3) + XCTAssertEqual(cache.offset, 7) } } diff --git a/packages/swift/Tests/NodeMLXCoreTests/MaskTests.swift b/packages/swift/Tests/NodeMLXCoreTests/MaskTests.swift new file mode 100644 index 0000000..7dd2b08 --- /dev/null +++ b/packages/swift/Tests/NodeMLXCoreTests/MaskTests.swift @@ -0,0 +1,129 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for attention mask creation. +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: tests/test_models.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import MLX +import XCTest + +@testable import NodeMLXCore + +final class MaskTests: XCTestCase { + // MARK: - Basic Causal Mask Tests + + func testBasicCausalMask() { + // A basic causal mask should be lower triangular + let mask = createCausalMask(n: 4, offset: 0) + + // Row 0: [T, F, F, F] + // Row 1: [T, T, F, F] + // Row 2: [T, T, T, F] + // Row 3: [T, T, T, T] + + XCTAssertEqual(mask.shape, [4, 4]) + + // First row: only first position visible + XCTAssertTrue(mask[0, 0].item(Bool.self)) + XCTAssertFalse(mask[0, 1].item(Bool.self)) + + // Last row: all positions visible + XCTAssertTrue(mask[3, 0].item(Bool.self)) + XCTAssertTrue(mask[3, 3].item(Bool.self)) + } + + func testCausalMaskWithOffset() { + // With offset, the mask should account for cached keys + let mask = createCausalMask(n: 3, offset: 2) + + // Shape should be [3, 5] (3 queries, 5 keys = 2 cached + 3 new) + XCTAssertEqual(mask.shape, [3, 5]) + + // First query can see all 3 keys (positions 0, 1, 2) + XCTAssertTrue(mask[0, 0].item(Bool.self)) + XCTAssertTrue(mask[0, 1].item(Bool.self)) + XCTAssertTrue(mask[0, 2].item(Bool.self)) + XCTAssertFalse(mask[0, 3].item(Bool.self)) + XCTAssertFalse(mask[0, 4].item(Bool.self)) + } + + // MARK: - Window Mask Tests + + func testMaskWithWindow() { + // Test sliding window attention mask + let mask = createCausalMask(n: 5, offset: 0, windowSize: 3) + + // With window size 3, each position can see at most 3 positions + // Row 0: [T, F, F, F, F] -> sum = 1 + // Row 1: [T, T, F, F, F] -> sum = 2 + // Row 2: [T, T, T, F, F] -> sum = 3 + // Row 3: [F, T, T, T, F] -> sum = 3 + // Row 4: [F, F, T, T, T] -> sum = 3 + + let expectedSums = [1, 2, 3, 3, 3] + for (i, expected) in expectedSums.enumerated() { + let rowSum = mask[i].asType(.int32).sum().item(Int.self) + XCTAssertEqual(rowSum, expected, "Row \(i) should have \(expected) visible positions") + } + } + + func testMaskWithWindowAndOffset() { + // Test sliding window with offset + let mask = createCausalMask(n: 5, offset: 1, windowSize: 3) + + // Shape should be [5, 6] (5 queries, 1 cached + 5 new) + XCTAssertEqual(mask.shape, [5, 6]) + + // First query at offset 1 can see positions 0 and 1 (within window) + // Expected sums: [2, 3, 3, 3, 3] + let expectedSums = [2, 3, 3, 3, 3] + for (i, expected) in expectedSums.enumerated() { + let rowSum = mask[i].asType(.int32).sum().item(Int.self) + XCTAssertEqual(rowSum, expected, "Row \(i) should have \(expected) visible positions") + } + } + + func testMaskWithWindowLargerOffset() { + // With larger offset, window should be fully utilized + let mask = createCausalMask(n: 5, offset: 2, windowSize: 3) + + // Shape: [5, 7] + XCTAssertEqual(mask.shape, [5, 7]) + + // All positions should see exactly 3 keys (window is full) + let expectedSums = [3, 3, 3, 3, 3] + for (i, expected) in expectedSums.enumerated() { + let rowSum = mask[i].asType(.int32).sum().item(Int.self) + XCTAssertEqual(rowSum, expected, "Row \(i) should have \(expected) visible positions") + } + } + + // MARK: - Edge Cases + + func testSingleTokenMask() { + let mask = createCausalMask(n: 1, offset: 0) + XCTAssertEqual(mask.shape, [1, 1]) + XCTAssertTrue(mask[0, 0].item(Bool.self)) + } + + func testSingleTokenWithOffset() { + let mask = createCausalMask(n: 1, offset: 5) + XCTAssertEqual(mask.shape, [1, 6]) + // Single query can see all 6 positions + let rowSum = mask[0].asType(.int32).sum().item(Int.self) + XCTAssertEqual(rowSum, 6) + } + + func testWindowSizeOne() { + let mask = createCausalMask(n: 4, offset: 0, windowSize: 1) + + // Each position can only see itself + for i in 0 ..< 4 { + let rowSum = mask[i].asType(.int32).sum().item(Int.self) + XCTAssertEqual(rowSum, 1, "Row \(i) should only see 1 position") + } + } +} diff --git a/packages/swift/Tests/NodeMLXCoreTests/ModelEvalTests.swift b/packages/swift/Tests/NodeMLXCoreTests/ModelEvalTests.swift deleted file mode 100644 index b5e4e70..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/ModelEvalTests.swift +++ /dev/null @@ -1,222 +0,0 @@ -// Copyright © 2026 Sebastian Software GmbH. -// Tests adapted from mlx-swift-lm patterns (MIT License, Apple Inc.) - -import MLX -import MLXNN -@testable import NodeMLXCore -import XCTest - -final class ModelEvalTests: XCTestCase { - // MARK: - Qwen2 Model Tests - - func testQwen2ModelForward() throws { - // Create a small Qwen2 model for testing - let config = try makeTestQwen2Config() - let model = Qwen2Model(config) - - // Quantize the model (like production usage) - quantize(model: model, groupSize: 64, bits: 4) - - // Forward pass with batch of tokens - let input = MLXArray([1, 2, 3, 4, 5])[.newAxis, .ellipsis] // [1, 5] - var cache: [KVCache]? = nil - let output = model(input, cache: &cache) - eval(output) - - XCTAssertEqual(output.shape, [1, 5, 100]) // [batch, seq, vocab] - } - - func testQwen2ModelWithCache() throws { - let config = try makeTestQwen2Config() - let model = Qwen2Model(config) - quantize(model: model, groupSize: 64, bits: 4) - - // Create cache - var cache: [KVCache]? = model.newCache() - XCTAssertEqual(cache?.count, 2) // One per layer - - // Prefill with 5 tokens - let prefill = MLXArray([1, 2, 3, 4, 5])[.newAxis, .ellipsis] - let out1 = model(prefill, cache: &cache) - eval(out1) - - XCTAssertEqual(out1.shape, [1, 5, 100]) - XCTAssertEqual(cache?[0].offset, 5) - - // Incremental generation - single token - let nextToken = MLXArray([6])[.newAxis, .ellipsis] - let out2 = model(nextToken, cache: &cache) - eval(out2) - - XCTAssertEqual(out2.shape, [1, 1, 100]) - XCTAssertEqual(cache?[0].offset, 6) - } - - // MARK: - Concurrent Evaluation Tests - - func testConcurrentModelEvaluation() async throws { - let config = try makeTestQwen2Config( - hiddenSize: 32, - intermediateSize: 64, - vocabSize: 50 - ) - let model = Qwen2Model(config) - quantize(model: model, groupSize: 64, bits: 4) - - // Force evaluation of all model weights before concurrent usage - eval(model) - - let numTasks = 3 - let results = await withTaskGroup(of: [Int].self) { group in - var allResults: [[Int]] = [] - - for taskId in 0 ..< numTasks { - group.addTask { - let input = MLXArray([1 + taskId, 2 + taskId, 3 + taskId])[.newAxis, .ellipsis] - let output = model(input) - eval(output) - return output.shape - } - } - - for await result in group { - allResults.append(result) - } - - return allResults - } - - XCTAssertEqual(results.count, numTasks) - - for result in results { - XCTAssertEqual(result, [1, 3, 50]) - } - } - - // MARK: - Sampling Tests - - func testGreedySampling() throws { - // Test that greedy sampling (argmax) works correctly - let vocabSize = 100 - let logits = MLXArray.zeros([vocabSize]) - - // Set one token to have highest probability - var logitsArray = logits.asArray(Float.self) - logitsArray[42] = 10.0 // Token 42 should be selected - let modifiedLogits = MLXArray(logitsArray) - - let token = argMax(modifiedLogits, axis: -1) - eval(token) - - XCTAssertEqual(Int(token.item(Int32.self)), 42) - } - - func testCategoricalSampling() throws { - // Test that categorical sampling produces valid tokens - let vocabSize = 100 - let logits = MLXRandom.normal([vocabSize]) - - // Sample multiple times and verify all tokens are in valid range - for _ in 0 ..< 10 { - let probs = softmax(logits, axis: -1) - let token = MLXRandom.categorical(probs) - eval(token) - - let tokenValue = Int(token.item(Int32.self)) - XCTAssertGreaterThanOrEqual(tokenValue, 0) - XCTAssertLessThan(tokenValue, vocabSize) - } - } - - // MARK: - Configuration Tests - - func testQwen2ConfigurationDecoding() throws { - let json = """ - { - "hidden_size": 1024, - "num_hidden_layers": 24, - "intermediate_size": 4096, - "num_attention_heads": 16, - "rms_norm_eps": 1e-6, - "vocab_size": 32000, - "num_key_value_heads": 4, - "rope_theta": 1000000.0 - } - """ - - let config = try JSONDecoder().decode( - Qwen2Configuration.self, - from: json.data(using: .utf8)! - ) - - XCTAssertEqual(config.hiddenSize, 1024) - XCTAssertEqual(config.numHiddenLayers, 24) - XCTAssertEqual(config.intermediateSize, 4096) - XCTAssertEqual(config.numAttentionHeads, 16) - XCTAssertEqual(config.vocabSize, 32000) - XCTAssertEqual(config.numKeyValueHeads, 4) - XCTAssertEqual(config.ropeTheta, 1_000_000.0) - } - - func testQwen2ConfigurationWithRopeScaling() throws { - let json = """ - { - "hidden_size": 512, - "num_hidden_layers": 8, - "intermediate_size": 2048, - "num_attention_heads": 8, - "rms_norm_eps": 1e-6, - "vocab_size": 10000, - "num_key_value_heads": 2, - "rope_theta": 10000.0, - "rope_scaling": { - "type": "linear", - "factor": 2.0 - } - } - """ - - let config = try JSONDecoder().decode( - Qwen2Configuration.self, - from: json.data(using: .utf8)! - ) - - XCTAssertNotNil(config.ropeScaling) - - if case let .string(type) = config.ropeScaling?["type"] { - XCTAssertEqual(type, "linear") - } else { - XCTFail("Expected rope_scaling type to be 'linear'") - } - - XCTAssertEqual(config.ropeScaling?["factor"]?.asFloat(), 2.0) - } -} - -// MARK: - Helper Functions for Testing - -/// Create a Qwen2Configuration from parameters (for testing) -func makeTestQwen2Config( - hiddenSize: Int = 64, - numHiddenLayers: Int = 2, - intermediateSize: Int = 128, - numAttentionHeads: Int = 4, - rmsNormEps: Float = 1e-6, - vocabSize: Int = 100, - numKeyValueHeads: Int = 2, - ropeTheta: Float = 10000.0 -) throws -> Qwen2Configuration { - let json = """ - { - "hidden_size": \(hiddenSize), - "num_hidden_layers": \(numHiddenLayers), - "intermediate_size": \(intermediateSize), - "num_attention_heads": \(numAttentionHeads), - "rms_norm_eps": \(rmsNormEps), - "vocab_size": \(vocabSize), - "num_key_value_heads": \(numKeyValueHeads), - "rope_theta": \(ropeTheta) - } - """ - return try JSONDecoder().decode(Qwen2Configuration.self, from: json.data(using: .utf8)!) -} diff --git a/packages/swift/Tests/NodeMLXCoreTests/ModelLoaderTests.swift b/packages/swift/Tests/NodeMLXCoreTests/ModelLoaderTests.swift deleted file mode 100644 index 568eb1e..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/ModelLoaderTests.swift +++ /dev/null @@ -1,94 +0,0 @@ -import Hub -import MLX -@testable import NodeMLXCore -import XCTest - -final class ModelLoaderTests: XCTestCase { - let loader = ModelLoader() - - // MARK: - Download Tests - - func testDownloadSmallModel() async throws { - // Use a small model to keep test fast - let modelId = "mlx-community/SmolLM-135M-4bit" - - print("Downloading \(modelId)...") - let modelDir = try await loader.download(modelId: modelId) { progress in - print(" Progress: \(Int(progress.fractionCompleted * 100))%") - } - - XCTAssertTrue(FileManager.default.fileExists(atPath: modelDir.path)) - print("✓ Downloaded to: \(modelDir.path)") - } - - // MARK: - Config Loading Tests - - func testLoadConfig() async throws { - let modelId = "mlx-community/SmolLM-135M-4bit" - let modelDir = try await loader.download(modelId: modelId) - - let config = try loader.loadConfig(from: modelDir) - - print("Model config:") - print(" model_type: \(config.modelType ?? "unknown")") - print(" hidden_size: \(config.hiddenSize ?? 0)") - print(" num_layers: \(config.numHiddenLayers ?? 0)") - print(" vocab_size: \(config.vocabSize ?? 0)") - - XCTAssertNotNil(config.modelType) - XCTAssertNotNil(config.hiddenSize) - } - - func testGetModelType() async throws { - let modelId = "mlx-community/SmolLM-135M-4bit" - let modelDir = try await loader.download(modelId: modelId) - - let modelType = try loader.getModelType(from: modelDir) - print("✓ Model type: \(modelType)") - - XCTAssertFalse(modelType.isEmpty) - } - - // MARK: - Weight Loading Tests - - func testLoadWeights() async throws { - let modelId = "mlx-community/SmolLM-135M-4bit" - let modelDir = try await loader.download(modelId: modelId) - - let weights = try loader.loadWeights(from: modelDir) - - print("Loaded \(weights.count) weight tensors:") - for (key, value) in weights.prefix(5) { - print(" \(key): \(value.shape)") - } - - XCTAssertFalse(weights.isEmpty) - print("✓ Successfully loaded \(weights.count) tensors") - } - - // MARK: - Weight Sanitization Tests - - func testSanitizeWeights() async throws { - let modelId = "mlx-community/SmolLM-135M-4bit" - let modelDir = try await loader.download(modelId: modelId) - - let rawWeights = try loader.loadWeights(from: modelDir) - let sanitized = sanitizeWeights(rawWeights, prefix: "model.") - - // Check if prefixes were removed - var prefixRemoved = false - for key in rawWeights.keys { - if key.hasPrefix("model.") { - let newKey = String(key.dropFirst("model.".count)) - if sanitized[newKey] != nil { - prefixRemoved = true - break - } - } - } - - print("✓ Weight sanitization completed") - print(" Original keys: \(rawWeights.count)") - print(" Sanitized keys: \(sanitized.count)") - } -} diff --git a/packages/swift/Tests/NodeMLXCoreTests/PerformanceTests.swift b/packages/swift/Tests/NodeMLXCoreTests/PerformanceTests.swift index 1470b82..b1a431d 100644 --- a/packages/swift/Tests/NodeMLXCoreTests/PerformanceTests.swift +++ b/packages/swift/Tests/NodeMLXCoreTests/PerformanceTests.swift @@ -1,7 +1,7 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT // -// PerformanceTests.swift -// Test MLXFast performance -// +// Performance tests for MLX operations. import MLX import MLXFast @@ -15,7 +15,6 @@ final class PerformanceTests: XCTestCase { let k = MLXArray.ones([1, 4, 8, 64]) let v = MLXArray.ones([1, 4, 8, 64]) - print("Testing MLXFast.scaledDotProductAttention...") let start = Date() let result = MLXFast.scaledDotProductAttention( queries: q, keys: k, values: v, @@ -25,9 +24,6 @@ final class PerformanceTests: XCTestCase { eval(result) let elapsed = Date().timeIntervalSince(start) - print("Shape: \(result.shape)") - print("Time: \(elapsed) s") - XCTAssertEqual(result.shape, [1, 4, 8, 64]) XCTAssertLessThan(elapsed, 1.0, "Attention should be fast!") } @@ -35,7 +31,6 @@ final class PerformanceTests: XCTestCase { func testRoPE() throws { let x = MLXArray.ones([1, 4, 8, 64]) - print("Testing MLXFast.RoPE...") let start = Date() let result = MLXFast.RoPE( x, @@ -48,30 +43,51 @@ final class PerformanceTests: XCTestCase { eval(result) let elapsed = Date().timeIntervalSince(start) - print("Shape: \(result.shape)") - print("Time: \(elapsed) s") - XCTAssertEqual(result.shape, [1, 4, 8, 64]) XCTAssertLessThan(elapsed, 0.5, "RoPE should be fast!") } func testKVCache() throws { - var cache = KVCacheSimple() + let cache = StandardKVCache() // First update let k1 = MLXArray.ones([1, 4, 8, 64]) let v1 = MLXArray.ones([1, 4, 8, 64]) - let (ck1, cv1) = cache.update(keys: k1, values: v1) + let (ck1, _) = cache.update(keys: k1, values: v1) XCTAssertEqual(ck1.dim(2), 8, "Cache should have 8 positions") XCTAssertEqual(cache.offset, 8) // Second update (single token) let k2 = MLXArray.ones([1, 4, 1, 64]) let v2 = MLXArray.ones([1, 4, 1, 64]) - let (ck2, cv2) = cache.update(keys: k2, values: v2) + let (ck2, _) = cache.update(keys: k2, values: v2) XCTAssertEqual(ck2.dim(2), 9, "Cache should have 9 positions") XCTAssertEqual(cache.offset, 9) + } + + func testMatmulPerformance() throws { + // Test basic matmul performance + let a = MLXArray.ones([256, 512]) + let b = MLXArray.ones([512, 256]) + + let start = Date() + let result = matmul(a, b) + eval(result) + let elapsed = Date().timeIntervalSince(start) + + XCTAssertEqual(result.shape, [256, 256]) + XCTAssertLessThan(elapsed, 0.5, "Matmul should be fast!") + } + + func testSoftmaxPerformance() throws { + let x = MLXArray.ones([1, 32, 128, 128]) + + let start = Date() + let result = softmax(x, axis: -1) + eval(result) + let elapsed = Date().timeIntervalSince(start) - print("KV Cache works correctly!") + XCTAssertEqual(result.shape, [1, 32, 128, 128]) + XCTAssertLessThan(elapsed, 0.5, "Softmax should be fast!") } } diff --git a/packages/swift/Tests/NodeMLXCoreTests/QuantizedKVCacheTests.swift b/packages/swift/Tests/NodeMLXCoreTests/QuantizedKVCacheTests.swift deleted file mode 100644 index aec96c3..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/QuantizedKVCacheTests.swift +++ /dev/null @@ -1,341 +0,0 @@ -// -// QuantizedKVCacheTests.swift -// NodeMLXCoreTests -// -// Additional tests for KVCache implementations - edge cases and advanced scenarios -// - -import MLX -import MLXFast -@testable import NodeMLXCore -import XCTest - -class AdditionalKVCacheTests: XCTestCase { - // MARK: - KVCacheSimple Edge Cases - - func testKVCacheSimpleLargeSequence() { - let cache = KVCacheSimple() - - // Test with sequence larger than default step size (256) - let keys = MLXArray.ones([1, 4, 300, 64]) - let values = MLXArray.ones([1, 4, 300, 64]) - - let (ck, cv) = cache.update(keys: keys, values: values) - - XCTAssertEqual(cache.offset, 300) - XCTAssertEqual(ck.dim(2), 300) - XCTAssertEqual(cv.dim(2), 300) - } - - func testKVCacheSimpleMultipleResets() { - let cache = KVCacheSimple() - - // First batch - let keys1 = MLXArray.ones([1, 4, 50, 64]) - let values1 = MLXArray.ones([1, 4, 50, 64]) - _ = cache.update(keys: keys1, values: values1) - XCTAssertEqual(cache.offset, 50) - - // Reset and start fresh - cache.reset() - XCTAssertEqual(cache.offset, 0) - - // New batch after reset - let keys2 = MLXArray.ones([1, 4, 100, 64]) - let values2 = MLXArray.ones([1, 4, 100, 64]) - let (ck2, _) = cache.update(keys: keys2, values: values2) - - XCTAssertEqual(cache.offset, 100) - XCTAssertEqual(ck2.dim(2), 100) - } - - func testKVCacheSimpleDifferentHeadDims() { - let cache = KVCacheSimple() - - // Keys and values can have different head dimensions - let keys = MLXArray.ones([1, 8, 16, 64]) // 8 heads, dim 64 - let values = MLXArray.ones([1, 8, 16, 128]) // 8 heads, dim 128 - - let (ck, cv) = cache.update(keys: keys, values: values) - - XCTAssertEqual(ck.dim(3), 64, "Key dimension should be preserved") - XCTAssertEqual(cv.dim(3), 128, "Value dimension should be preserved") - } - - func testKVCacheSimpleWithBatchSize() { - let cache = KVCacheSimple() - - // Test with batch size > 1 - let keys = MLXArray.ones([4, 8, 16, 64]) // batch=4 - let values = MLXArray.ones([4, 8, 16, 64]) - - let (ck, cv) = cache.update(keys: keys, values: values) - - XCTAssertEqual(ck.dim(0), 4, "Batch dimension should be preserved") - XCTAssertEqual(cv.dim(0), 4) - } - - // MARK: - RotatingKVCache Edge Cases - - func testRotatingKVCacheKeepParameter() { - // Test that 'keep' tokens are preserved during rotation - let cache = RotatingKVCache(maxSize: 100, keep: 10, step: 50) - - // Fill cache past rotation point - for _ in 0 ..< 3 { - let keys = MLXArray.ones([1, 4, 50, 64]) - let values = MLXArray.ones([1, 4, 50, 64]) - _ = cache.update(keys: keys, values: values) - } - - // Offset should track total tokens seen - XCTAssertEqual(cache.offset, 150) - } - - func testRotatingKVCacheSingleTokenUpdates() { - let cache = RotatingKVCache(maxSize: 100, keep: 0) - - // Initial batch - let initKeys = MLXArray.ones([1, 4, 50, 64]) - let initValues = MLXArray.ones([1, 4, 50, 64]) - _ = cache.update(keys: initKeys, values: initValues) - - // Single token updates (typical for generation) - for i in 0 ..< 60 { - let keys = MLXArray.ones([1, 4, 1, 64]) - let values = MLXArray.ones([1, 4, 1, 64]) - let (ck, _) = cache.update(keys: keys, values: values) - - if cache.offset <= 100 { - XCTAssertEqual(ck.dim(2), 51 + i, "Cache should grow until maxSize") - } - } - - // Final offset - XCTAssertEqual(cache.offset, 110) - } - - func testRotatingKVCacheMaxSizeProperty() { - let cache = RotatingKVCache(maxSize: 512) - XCTAssertEqual(cache.maxSize, 512) - } - - // MARK: - MakeMask Tests - - func testKVCacheSimpleMakeMaskSingleToken() { - let cache = KVCacheSimple() - - // Add some data to the cache - let keys = MLXArray.ones([1, 4, 10, 64]) - let values = MLXArray.ones([1, 4, 10, 64]) - _ = cache.update(keys: keys, values: values) - - // Single token should return no mask - let mask = cache.makeMask(n: 1, windowSize: nil, returnArray: false) - - switch mask { - case .none: - break // Expected - default: - XCTFail("Single token should return .none mask") - } - } - - func testKVCacheSimpleMakeMaskMultiToken() { - let cache = KVCacheSimple() - - let mask = cache.makeMask(n: 10, windowSize: nil, returnArray: false) - - switch mask { - case .causal: - break // Expected for multi-token without window - default: - XCTFail("Multi-token should return .causal mask") - } - } - - func testKVCacheSimpleMakeMaskWithWindow() { - let cache = KVCacheSimple() - - // Add data - let keys = MLXArray.ones([1, 4, 50, 64]) - let values = MLXArray.ones([1, 4, 50, 64]) - _ = cache.update(keys: keys, values: values) - - // Request mask with window size smaller than sequence - let mask = cache.makeMask(n: 20, windowSize: 10, returnArray: true) - - switch mask { - case let .array(arr): - // Should have the mask array - XCTAssertGreaterThan(arr.size, 0) - default: - XCTFail("Should return array mask when window size specified") - } - } - - func testRotatingKVCacheMakeMaskAfterRotation() { - let cache = RotatingKVCache(maxSize: 50, keep: 5) - - // Fill past rotation - let keys = MLXArray.ones([1, 4, 100, 64]) - let values = MLXArray.ones([1, 4, 100, 64]) - _ = cache.update(keys: keys, values: values) - - // Mask after rotation - let mask = cache.makeMask(n: 10, windowSize: 30, returnArray: true) - - switch mask { - case .array: - break // Expected - case .causal: - break // Also acceptable - default: - XCTFail("Should return array or causal mask") - } - } - - // MARK: - createLayerCaches Tests - - func testCreateLayerCachesDefaultType() { - let caches = createLayerCaches(numLayers: 32) - - XCTAssertEqual(caches.count, 32) - XCTAssertTrue(caches[0] is KVCacheSimple) - XCTAssertTrue(caches[31] is KVCacheSimple) - } - - func testCreateLayerCachesWithMaxKVSize() { - let caches = createLayerCaches(numLayers: 24, maxKVSize: 4096) - - XCTAssertEqual(caches.count, 24) - - // All should be RotatingKVCache - for cache in caches { - XCTAssertTrue(cache is RotatingKVCache, "Should create RotatingKVCache when maxKVSize specified") - } - } - - // MARK: - createCausalMask Tests - - func testCreateCausalMaskBasic() { - let mask = createCausalMask(n: 4, offset: 0) - eval(mask) - - // Should be a lower triangular matrix - XCTAssertEqual(mask.shape, [4, 4]) - - // Check diagonal and below are true - XCTAssertTrue(mask[0, 0].item(Bool.self)) - XCTAssertTrue(mask[1, 1].item(Bool.self)) - XCTAssertTrue(mask[2, 2].item(Bool.self)) - XCTAssertTrue(mask[3, 3].item(Bool.self)) - - // Check above diagonal is false - XCTAssertFalse(mask[0, 1].item(Bool.self)) - XCTAssertFalse(mask[0, 3].item(Bool.self)) - } - - func testCreateCausalMaskWithOffset() { - let mask = createCausalMask(n: 2, offset: 3) - eval(mask) - - // With offset 3, new tokens (positions 3,4) can attend to old (0,1,2) and themselves - XCTAssertEqual(mask.shape, [2, 5]) // 2 new tokens, 5 total positions - - // First new token (pos 3) can see positions 0-3 - XCTAssertTrue(mask[0, 0].item(Bool.self)) - XCTAssertTrue(mask[0, 3].item(Bool.self)) - XCTAssertFalse(mask[0, 4].item(Bool.self)) // Can't see future - } - - func testCreateCausalMaskWithWindowSize() { - let mask = createCausalMask(n: 4, offset: 0, windowSize: 2) - eval(mask) - - XCTAssertEqual(mask.shape, [4, 4]) - - // With window size 2, each position can only see 2 previous positions - // Position 3 can see positions 2 and 3 (not 0 and 1) - XCTAssertFalse(mask[3, 0].item(Bool.self), "Position 3 should not see position 0 with window=2") - XCTAssertFalse(mask[3, 1].item(Bool.self), "Position 3 should not see position 1 with window=2") - XCTAssertTrue(mask[3, 2].item(Bool.self), "Position 3 should see position 2") - XCTAssertTrue(mask[3, 3].item(Bool.self), "Position 3 should see itself") - } - - // MARK: - createAttentionMask Function Tests - - func testCreateAttentionMaskNilCache() { - let h = MLXArray.ones([1, 10, 64]) // seq_len = 10 - - let mask = createAttentionMask(h: h, cache: nil, windowSize: nil, returnArray: false) - - switch mask { - case .causal: - break // Expected for seq > 1 - default: - XCTFail("Should return causal mask for multi-token without cache") - } - } - - func testCreateAttentionMaskSingleToken() { - let h = MLXArray.ones([1, 1, 64]) // seq_len = 1 - - let mask = createAttentionMask(h: h, cache: nil, windowSize: nil, returnArray: false) - - switch mask { - case .none: - break // Expected for single token - default: - XCTFail("Should return .none mask for single token") - } - } - - func testCreateAttentionMaskWithCache() { - let cache = KVCacheSimple() - - // Add some data - let keys = MLXArray.ones([1, 4, 20, 64]) - let values = MLXArray.ones([1, 4, 20, 64]) - _ = cache.update(keys: keys, values: values) - - let h = MLXArray.ones([1, 5, 64]) // New 5 tokens - - let mask = createAttentionMask(h: h, cache: cache, windowSize: nil, returnArray: false) - - // Should delegate to cache.makeMask - switch mask { - case .causal: - break // Expected - default: - XCTFail("Should return causal mask from cache") - } - } - - // MARK: - Data Type Preservation Tests - - func testKVCachePreservesDtype() { - let cache = KVCacheSimple() - - // Test with float16 - let keys = MLXArray.ones([1, 4, 10, 64]).asType(.float16) - let values = MLXArray.ones([1, 4, 10, 64]).asType(.float16) - - let (ck, cv) = cache.update(keys: keys, values: values) - - XCTAssertEqual(ck.dtype, .float16, "Cache should preserve key dtype") - XCTAssertEqual(cv.dtype, .float16, "Cache should preserve value dtype") - } - - func testKVCacheWithBFloat16() { - let cache = KVCacheSimple() - - let keys = MLXArray.ones([1, 4, 10, 64]).asType(.bfloat16) - let values = MLXArray.ones([1, 4, 10, 64]).asType(.bfloat16) - - let (ck, cv) = cache.update(keys: keys, values: values) - - XCTAssertEqual(ck.dtype, .bfloat16) - XCTAssertEqual(cv.dtype, .bfloat16) - } -} diff --git a/packages/swift/Tests/NodeMLXCoreTests/RoPETests.swift b/packages/swift/Tests/NodeMLXCoreTests/RoPETests.swift deleted file mode 100644 index 4e376fa..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/RoPETests.swift +++ /dev/null @@ -1,320 +0,0 @@ -// -// RoPETests.swift -// NodeMLXCoreTests -// -// Tests for Rotary Position Embedding implementations -// - -import MLX -import MLXNN -@testable import NodeMLXCore -import XCTest - -class RoPETests: XCTestCase { - // MARK: - initializeRope Factory Tests - - func testInitializeRopeDefault() { - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: nil, - maxPositionEmbeddings: 2048 - ) - - XCTAssertTrue(rope is RoPE, "Default should create standard RoPE") - } - - func testInitializeRopeLinear() { - let config: [String: StringOrNumber] = [ - "type": .string("linear"), - "factor": .float(2.0), - ] - - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: config, - maxPositionEmbeddings: 2048 - ) - - XCTAssertTrue(rope is RoPE, "Linear should create standard RoPE with scale") - } - - func testInitializeRopeLlama3() { - let config: [String: StringOrNumber] = [ - "type": .string("llama3"), - "factor": .float(8.0), - "low_freq_factor": .float(1.0), - "high_freq_factor": .float(4.0), - "original_max_position_embeddings": .int(8192), - ] - - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: config, - maxPositionEmbeddings: 131_072 - ) - - XCTAssertTrue(rope is Llama3RoPE, "llama3 type should create Llama3RoPE") - } - - func testInitializeRopeYarn() { - let config: [String: StringOrNumber] = [ - "type": .string("yarn"), - "factor": .float(16.0), - "original_max_position_embeddings": .int(4096), - "beta_fast": .float(32.0), - "beta_slow": .float(1.0), - "mscale": .float(1.0), - "mscale_all_dim": .float(0.0), - ] - - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: config, - maxPositionEmbeddings: 65536 - ) - - XCTAssertTrue(rope is YarnRoPE, "yarn type should create YarnRoPE") - } - - func testInitializeRopeLongrope() { - // LongRoPE requires short_factor and long_factor arrays - let config: [String: StringOrNumber] = [ - "type": .string("longrope"), - "original_max_position_embeddings": .int(4096), - "short_factor": .floats(Array(repeating: 1.0, count: 32)), - "long_factor": .floats(Array(repeating: 2.0, count: 32)), - ] - - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: config, - maxPositionEmbeddings: 131_072 - ) - - XCTAssertTrue(rope is SuScaledRoPE, "longrope type should create SuScaledRoPE") - } - - func testInitializeRopeMrope() { - let config: [String: StringOrNumber] = [ - "type": .string("mrope"), - "mrope_section": .ints([16, 24, 24]), - ] - - let rope = initializeRope( - dims: 64, - base: 10000.0, - traditional: false, - scalingConfig: config, - maxPositionEmbeddings: 2048 - ) - - // MRoPE falls back to standard RoPE - XCTAssertTrue(rope is RoPE, "mrope type should create standard RoPE") - } - - // MARK: - Llama3RoPE Tests - - func testLlama3RoPEOutput() { - let config: [String: StringOrNumber] = [ - "factor": .float(8.0), - "low_freq_factor": .float(1.0), - "high_freq_factor": .float(4.0), - "original_max_position_embeddings": .int(8192), - ] - - let rope = Llama3RoPE( - dims: 64, - maxPositionEmbeddings: 131_072, - traditional: false, - base: 10000.0, - scalingConfig: config - ) - - // Test input: [batch=1, seq=4, heads=8, dims=64] - let input = MLXArray.ones([1, 4, 8, 64]) - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testLlama3RoPEWithOffset() { - let config: [String: StringOrNumber] = [ - "factor": .float(8.0), - "low_freq_factor": .float(1.0), - "high_freq_factor": .float(4.0), - "original_max_position_embeddings": .int(8192), - ] - - let rope = Llama3RoPE( - dims: 64, - maxPositionEmbeddings: 131_072, - traditional: false, - base: 10000.0, - scalingConfig: config - ) - - let input = MLXArray.ones([1, 1, 8, 64]) - - // Same input at different offsets should produce different outputs - let output0 = rope(input, offset: 0) - let output100 = rope(input, offset: 100) - eval(output0, output100) - - // Check that outputs differ - let diff = abs(output0 - output100) - let maxDiff = MLX.max(diff).item(Float.self) - XCTAssertGreaterThan(maxDiff, 0.01, "Different offsets should produce different embeddings") - } - - // MARK: - YarnRoPE Tests - - func testYarnRoPEOutput() { - let rope = YarnRoPE( - dimensions: 64, - traditional: false, - maxPositionEmbeddings: 65536, - base: 10000.0, - scalingFactor: 16.0, - originalMaxPositionEmbeddings: 4096, - betaFast: 32.0, - betaSlow: 1.0, - mscale: 1.0, - mscaleAllDim: 0.0 - ) - - // Test input - let input = MLXArray.ones([1, 4, 8, 64]) - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testYarnRoPEWithMscale() { - // Test with mscale != 1.0 - let rope = YarnRoPE( - dimensions: 64, - traditional: false, - maxPositionEmbeddings: 65536, - base: 10000.0, - scalingFactor: 16.0, // > 1 will activate mscale - originalMaxPositionEmbeddings: 4096, - betaFast: 32.0, - betaSlow: 1.0, - mscale: 0.707, // Custom mscale - mscaleAllDim: 0.0 - ) - - let input = MLXArray.ones([1, 4, 8, 64]) - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - // MARK: - SuScaledRoPE Tests - - func testSuScaledRoPEShortContext() { - let rope = SuScaledRoPE( - dimensions: 64, - base: 10000.0, - maxPositionEmbeddings: 131_072, - originalMaxPositionEmbeddings: 4096, - shortFactor: Array(repeating: 1.0, count: 32), - longFactor: Array(repeating: 2.0, count: 32) - ) - - // Short context (within original max) - let input = MLXArray.ones([1, 100, 8, 64]) // seq=100 < 4096 - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testSuScaledRoPELongContext() { - let rope = SuScaledRoPE( - dimensions: 64, - base: 10000.0, - maxPositionEmbeddings: 131_072, - originalMaxPositionEmbeddings: 4096, - shortFactor: Array(repeating: 1.0, count: 32), - longFactor: Array(repeating: 2.0, count: 32) - ) - - // Long context (beyond original max using offset) - let input = MLXArray.ones([1, 100, 8, 64]) - let output = rope(input, offset: 5000) // 100 + 5000 > 4096 - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testSuScaledRoPEDifferentContextLengths() { - let rope = SuScaledRoPE( - dimensions: 64, - base: 10000.0, - maxPositionEmbeddings: 131_072, - originalMaxPositionEmbeddings: 4096, - shortFactor: Array(repeating: 1.0, count: 32), - longFactor: Array(repeating: 3.0, count: 32) - ) - - let input = MLXArray.ones([1, 100, 8, 64]) - - // Short vs long context should produce different outputs - let outputShort = rope(input, offset: 0) // seq_len = 100 < 4096 - let outputLong = rope(input, offset: 4000) // seq_len = 4100 > 4096 - eval(outputShort, outputLong) - - let diff = abs(outputShort - outputLong) - let maxDiff = MLX.max(diff).item(Float.self) - - // Should use different frequency factors - XCTAssertGreaterThan(maxDiff, 0.001, "Short and long context should produce different embeddings") - } - - // MARK: - Standard RoPE Reference Tests - - func testStandardRoPEBasic() { - let rope = RoPE(dimensions: 64, traditional: false, base: 10000.0, scale: 1.0) - - let input = MLXArray.ones([1, 4, 8, 64]) - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testStandardRoPEWithScale() { - let rope = RoPE(dimensions: 64, traditional: false, base: 10000.0, scale: 0.5) - - let input = MLXArray.ones([1, 4, 8, 64]) - let output = rope(input, offset: 0) - - XCTAssertEqual(output.shape, input.shape, "Output shape should match input shape") - } - - func testRoPETraditionalMode() { - // Traditional mode uses different rotation formula - let ropeTraditional = RoPE(dimensions: 64, traditional: true, base: 10000.0, scale: 1.0) - let ropeModern = RoPE(dimensions: 64, traditional: false, base: 10000.0, scale: 1.0) - - let input = MLXArray.ones([1, 4, 8, 64]) * 0.5 // Non-trivial values - - let outputTraditional = ropeTraditional(input, offset: 10) - let outputModern = ropeModern(input, offset: 10) - eval(outputTraditional, outputModern) - - // Traditional and modern modes should produce different results - let diff = abs(outputTraditional - outputModern) - let maxDiff = MLX.max(diff).item(Float.self) - - XCTAssertGreaterThan(maxDiff, 0.001, "Traditional and modern RoPE should differ") - } -} diff --git a/packages/swift/Tests/NodeMLXCoreTests/RoPEUtilsTests.swift b/packages/swift/Tests/NodeMLXCoreTests/RoPEUtilsTests.swift new file mode 100644 index 0000000..c5f947a --- /dev/null +++ b/packages/swift/Tests/NodeMLXCoreTests/RoPEUtilsTests.swift @@ -0,0 +1,259 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for ported/RoPEUtils.swift + +import MLX +import MLXNN +import XCTest + +@testable import NodeMLXCore + +final class RoPEUtilsTests: XCTestCase { + // MARK: - StandardRoPE Tests + + func testStandardRoPEInitialization() { + let rope = StandardRoPE(dims: 64) + XCTAssertNotNil(rope) + } + + func testStandardRoPEApply() { + let rope = StandardRoPE(dims: 64, base: 10000.0) + + // Create test input: [batch, heads, seq, dim] + let input = MLXArray.ones([1, 4, 8, 64]) + + // Apply RoPE at offset 0 + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } + + func testStandardRoPEWithOffset() { + let rope = StandardRoPE(dims: 64) + + let input = MLXArray.ones([1, 4, 1, 64]) + + // Apply at different offsets + let output1 = rope(input, offset: 0) + let output2 = rope(input, offset: 10) + + // Outputs should be different due to different positions + XCTAssertEqual(output1.shape, output2.shape) + // Note: We can't easily compare values, but shapes should match + } + + // MARK: - Llama3RoPE Tests + + func testLlama3RoPEInitialization() { + let config: [String: Any] = [ + "factor": 8.0, + "low_freq_factor": 1.0, + "high_freq_factor": 4.0, + "original_max_position_embeddings": 8192, + ] + let rope = Llama3RoPE( + dims: 64, + maxPositionEmbeddings: 8192, + base: 500_000.0, + scalingConfig: config + ) + XCTAssertNotNil(rope) + } + + func testLlama3RoPEApply() { + let config: [String: Any] = [ + "factor": 8.0, + "low_freq_factor": 1.0, + "high_freq_factor": 4.0, + "original_max_position_embeddings": 8192, + ] + let rope = Llama3RoPE( + dims: 64, + maxPositionEmbeddings: 8192, + base: 500_000.0, + scalingConfig: config + ) + + let input = MLXArray.ones([1, 4, 8, 64]) + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } + + // MARK: - SuScaledRoPE Tests + + func testSuScaledRoPEInitialization() { + let longFactor = [Float](repeating: 1.0, count: 32) + let rope = SuScaledRoPE( + dims: 64, + maxPositionEmbeddings: 131_072, + originalMaxPositionEmbeddings: 4096, + longFactor: longFactor + ) + XCTAssertNotNil(rope) + } + + func testSuScaledRoPEApply() { + // Create with proper long factor (one per dimension pair) + let longFactor = [Float](repeating: 1.0, count: 32) // 64 dims / 2 + let rope = SuScaledRoPE( + dims: 64, + maxPositionEmbeddings: 131_072, + originalMaxPositionEmbeddings: 4096, + longFactor: longFactor + ) + + let input = MLXArray.ones([1, 4, 8, 64]) + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } + + // MARK: - YarnRoPE Tests + + func testYarnRoPEInitialization() { + let rope = YarnRoPE( + dims: 64, + traditional: false, + base: 10000.0, + scalingFactor: 1.0 + ) + XCTAssertNotNil(rope) + } + + func testYarnRoPEApply() { + let rope = YarnRoPE( + dims: 64, + traditional: false, + base: 10000.0, + scalingFactor: 1.0 + ) + + let input = MLXArray.ones([1, 4, 8, 64]) + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } + + // MARK: - Factory Function Tests + + func testInitializeRopeDefault() { + let rope = initializeRope( + dims: 64, + base: 10000.0, + traditional: false, + scalingConfig: nil, + maxPositionEmbeddings: nil + ) + + XCTAssertTrue(rope is StandardRoPE) + } + + func testInitializeRopeLlama3() { + let config: [String: Any] = [ + "type": "llama3", + "factor": 8.0, + "low_freq_factor": 1.0, + "high_freq_factor": 4.0, + "original_max_position_embeddings": 8192, + ] + + let rope = initializeRope( + dims: 64, + base: 500_000.0, + traditional: false, + scalingConfig: config, + maxPositionEmbeddings: 131_072 + ) + + XCTAssertTrue(rope is Llama3RoPE) + } + + func testInitializeRopeYarn() { + let config: [String: Any] = [ + "type": "yarn", + "factor": 2.0, + "attention_factor": 1.0, + "beta_fast": 32.0, + "beta_slow": 1.0, + "original_max_position_embeddings": 4096, + ] + + let rope = initializeRope( + dims: 64, + base: 10000.0, + traditional: false, + scalingConfig: config, + maxPositionEmbeddings: 8192 + ) + + XCTAssertTrue(rope is YarnRoPE) + } + + func testInitializeRopeSuScaled() { + // Note: "su" type maps to "longrope" in initializeRope + let config: [String: Any] = [ + "type": "longrope", + "long_factor": [Float](repeating: 1.0, count: 32), + "original_max_position_embeddings": 4096, + ] + + let rope = initializeRope( + dims: 64, + base: 10000.0, + traditional: false, + scalingConfig: config, + maxPositionEmbeddings: 131_072 + ) + + XCTAssertTrue(rope is SuScaledRoPE) + } + + func testInitializeRopeLongRope() { + let config: [String: Any] = [ + "type": "longrope", + "long_factor": [Float](repeating: 1.0, count: 32), + "original_max_position_embeddings": 4096, + ] + + let rope = initializeRope( + dims: 64, + base: 10000.0, + traditional: false, + scalingConfig: config, + maxPositionEmbeddings: 131_072 + ) + + XCTAssertTrue(rope is SuScaledRoPE) + } + + // MARK: - Edge Cases + + func testRoPEWithSmallDimensions() { + let rope = StandardRoPE(dims: 8) + let input = MLXArray.ones([1, 1, 4, 8]) + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } + + func testRoPEWithLargeOffset() { + let rope = StandardRoPE(dims: 64) + let input = MLXArray.ones([1, 4, 1, 64]) + + // Large offset simulating long context + let output = rope(input, offset: 10000) + + XCTAssertEqual(output.shape, input.shape) + } + + func testRoPEWithBatchSize() { + let rope = StandardRoPE(dims: 64) + let input = MLXArray.ones([4, 8, 16, 64]) // batch=4 + + let output = rope(input, offset: 0) + + XCTAssertEqual(output.shape, input.shape) + } +} diff --git a/packages/swift/Tests/NodeMLXCoreTests/SamplingUtilsTests.swift b/packages/swift/Tests/NodeMLXCoreTests/SamplingUtilsTests.swift new file mode 100644 index 0000000..3d3845d --- /dev/null +++ b/packages/swift/Tests/NodeMLXCoreTests/SamplingUtilsTests.swift @@ -0,0 +1,203 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for SamplingUtils. +// +// Ported from mlx-lm (https://github.com/ml-explore/mlx-lm) +// Original: tests/test_sample_utils.py +// Git Hash: 7585c142a6be9c9245f4ce61d087839776cb8275 (2026-01-12) + +import MLX +import XCTest + +@testable import NodeMLXCore + +final class SamplingUtilsTests: XCTestCase { + // MARK: - Top-P Tests + + func testApplyTopPHighConfidence() { + // When top token has 0.9 probability and threshold is 0.3, + // only the top token should remain + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopP(logits, p: 0.3) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + XCTAssertEqual(actualProbs[0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[2].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.0, accuracy: 1e-5) + } + + func testApplyTopPHighThreshold() { + // When threshold is 0.95, all tokens should remain + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopP(logits, p: 0.95) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + XCTAssertEqual(actualProbs[0].item(Float.self), 0.9, accuracy: 1e-4) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.1, accuracy: 1e-4) + } + + func testApplyTopPMultipleTokens() { + let probs = MLXArray([Float(0.0), 0.5, 0.4, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + // With p=0.4, only the top token should remain + var newLogits = SamplingUtils.applyTopP(logits, p: 0.4) + var actualProbs = softmax(newLogits, axis: -1).squeezed() + XCTAssertEqual(actualProbs[0].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[2].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.0, accuracy: 1e-5) + + // With p=0.6, top two tokens should remain + newLogits = SamplingUtils.applyTopP(logits, p: 0.6) + actualProbs = softmax(newLogits, axis: -1).squeezed() + XCTAssertEqual(actualProbs[0].item(Float.self), 0.0, accuracy: 1e-4) + XCTAssertEqual(actualProbs[1].item(Float.self), 0.5556, accuracy: 1e-3) + XCTAssertEqual(actualProbs[2].item(Float.self), 0.4444, accuracy: 1e-3) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.0, accuracy: 1e-4) + } + + func testApplyTopPBatchMode() { + // Create 2x4 batch: [[0.9, 0.0, 0.0, 0.1], [0.0, 0.8, 0.1, 0.1]] + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1, 0.0, 0.8, 0.1, 0.1]).reshaped([2, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopP(logits, p: 0.5) + let actualProbs = softmax(newLogits, axis: -1) + + // First batch: only first token + XCTAssertEqual(actualProbs[0, 0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[0, 1].item(Float.self), 0.0, accuracy: 1e-5) + + // Second batch: only second token + XCTAssertEqual(actualProbs[1, 0].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1, 1].item(Float.self), 1.0, accuracy: 1e-5) + } + + // MARK: - Top-K Tests + + func testApplyTopKSingle() { + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopK(logits, k: 1) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + XCTAssertEqual(actualProbs[0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[2].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.0, accuracy: 1e-5) + } + + func testApplyTopKTwo() { + let probs = MLXArray([Float(0.6), 0.0, 0.1, 0.3]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopK(logits, k: 2) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + // Renormalized: 0.6/(0.6+0.3) = 0.6667, 0.3/(0.6+0.3) = 0.3333 + XCTAssertEqual(actualProbs[0].item(Float.self), 0.6667, accuracy: 1e-3) + XCTAssertEqual(actualProbs[1].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[2].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.3333, accuracy: 1e-3) + } + + func testApplyTopKBatchMode() { + // Create 2x4 batch: [[0.9, 0.0, 0.0, 0.1], [0.0, 0.8, 0.0, 0.1]] + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1, 0.0, 0.8, 0.0, 0.1]).reshaped([2, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyTopK(logits, k: 1) + let actualProbs = softmax(newLogits, axis: -1) + + // First batch: only first token + XCTAssertEqual(actualProbs[0, 0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[0, 3].item(Float.self), 0.0, accuracy: 1e-5) + + // Second batch: only second token + XCTAssertEqual(actualProbs[1, 0].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1, 1].item(Float.self), 1.0, accuracy: 1e-5) + } + + // MARK: - Min-P Tests + + func testApplyMinPHighThreshold() { + // With minP=0.8, only tokens with prob >= 0.8 * maxProb remain + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyMinP(logits, minP: 0.8) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + // Only first token (0.9) passes: 0.1 < 0.8 * 0.9 = 0.72 + XCTAssertEqual(actualProbs[0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.0, accuracy: 1e-5) + } + + func testApplyMinPLowThreshold() { + // With minP=0.05, tokens with prob >= 0.05 * maxProb remain + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + let newLogits = SamplingUtils.applyMinP(logits, minP: 0.05) + let actualProbs = softmax(newLogits, axis: -1).squeezed() + + // Both first and last pass: 0.1 >= 0.05 * 0.9 = 0.045 + XCTAssertEqual(actualProbs[0].item(Float.self), 0.9, accuracy: 1e-4) + XCTAssertEqual(actualProbs[3].item(Float.self), 0.1, accuracy: 1e-4) + } + + func testApplyMinPBatchMode() { + // Create 2x4 batch: [[0.9, 0.0, 0.0, 0.1], [0.0, 0.8, 0.0, 0.1]] + let probs = MLXArray([Float(0.9), 0.0, 0.0, 0.1, 0.0, 0.8, 0.0, 0.1]).reshaped([2, 4]) + let logits = log(probs) + + // With minP=0.7, threshold is 0.7 * maxProb + let newLogits = SamplingUtils.applyMinP(logits, minP: 0.7) + let actualProbs = softmax(newLogits, axis: -1) + + // First batch: threshold = 0.7 * 0.9 = 0.63, only first passes + XCTAssertEqual(actualProbs[0, 0].item(Float.self), 1.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[0, 3].item(Float.self), 0.0, accuracy: 1e-5) + + // Second batch: threshold = 0.7 * 0.8 = 0.56, only second passes + XCTAssertEqual(actualProbs[1, 0].item(Float.self), 0.0, accuracy: 1e-5) + XCTAssertEqual(actualProbs[1, 1].item(Float.self), 1.0, accuracy: 1e-5) + } + + // MARK: - Combined Sampling Tests + + func testSampleTokenGreedy() { + let probs = MLXArray([Float(0.1), 0.2, 0.5, 0.2]).reshaped([1, 4]) + let logits = log(probs) + + // With temperature 0, should always pick highest probability + let token = SamplingUtils.sampleToken(logits: logits, temperature: 0) + XCTAssertEqual(token, 2) // Index of 0.5 + } + + func testSampleTokenWithTopK() { + let probs = MLXArray([Float(0.1), 0.2, 0.5, 0.2]).reshaped([1, 4]) + let logits = log(probs) + + // With topK=1 and temp=0, should pick highest + let token = SamplingUtils.sampleToken(logits: logits, temperature: 0, topK: 1) + XCTAssertEqual(token, 2) + } + + func testSampleTokenWithTopP() { + let probs = MLXArray([Float(0.1), 0.2, 0.6, 0.1]).reshaped([1, 4]) + let logits = log(probs) + + // With topP=0.5 and temp=0, should pick highest (only token in nucleus) + let token = SamplingUtils.sampleToken(logits: logits, temperature: 0, topP: 0.5) + XCTAssertEqual(token, 2) + } +} diff --git a/packages/swift/Tests/NodeMLXCoreTests/StringOrNumberTests.swift b/packages/swift/Tests/NodeMLXCoreTests/StringOrNumberTests.swift index c9841b8..0b948dc 100644 --- a/packages/swift/Tests/NodeMLXCoreTests/StringOrNumberTests.swift +++ b/packages/swift/Tests/NodeMLXCoreTests/StringOrNumberTests.swift @@ -1,5 +1,7 @@ -// Copyright © 2026 Sebastian Software GmbH. -// Tests adapted from mlx-swift-lm patterns (MIT License, Apple Inc.) +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for StringOrNumber JSON type handling. import Foundation @testable import NodeMLXCore @@ -30,81 +32,41 @@ final class StringOrNumberTests: XCTestCase { } } - func testDecodeFloat() throws { + func testDecodeDouble() throws { let json = "3.14" let value = try JSONDecoder().decode(StringOrNumber.self, from: json.data(using: .utf8)!) - if case let .float(f) = value { - XCTAssertEqual(f, 3.14, accuracy: 0.001) + if case let .double(d) = value { + XCTAssertEqual(d, 3.14, accuracy: 0.001) } else { - XCTFail("Expected float") + XCTFail("Expected double") } } - func testDecodeBool() throws { - let json = "true" - let value = try JSONDecoder().decode(StringOrNumber.self, from: json.data(using: .utf8)!) - - if case let .bool(b) = value { - XCTAssertTrue(b) - } else { - XCTFail("Expected bool") - } - } - - func testDecodeIntArray() throws { - let json = "[1, 2, 3]" - let value = try JSONDecoder().decode(StringOrNumber.self, from: json.data(using: .utf8)!) - - if case let .ints(arr) = value { - XCTAssertEqual(arr, [1, 2, 3]) - } else { - XCTFail("Expected int array") - } - } + // MARK: - Value Accessor Tests - func testDecodeFloatArray() throws { - let json = "[1.1, 2.2, 3.3]" - let value = try JSONDecoder().decode(StringOrNumber.self, from: json.data(using: .utf8)!) - - if case let .floats(arr) = value { - XCTAssertEqual(arr.count, 3) - XCTAssertEqual(arr[0], 1.1, accuracy: 0.001) - } else { - XCTFail("Expected float array") - } + func testStringValue() throws { + XCTAssertEqual(StringOrNumber.string("hello").stringValue, "hello") + XCTAssertEqual(StringOrNumber.int(42).stringValue, "42") + XCTAssertEqual(StringOrNumber.double(3.14).stringValue, "3.14") } - // MARK: - Conversion Tests - - func testAsFloat() throws { - XCTAssertEqual(StringOrNumber.int(42).asFloat(), 42.0) - XCTAssertEqual(StringOrNumber.float(3.14).asFloat(), 3.14) - XCTAssertNil(StringOrNumber.string("hello").asFloat()) - XCTAssertEqual(StringOrNumber.bool(true).asFloat(), 1.0) - XCTAssertEqual(StringOrNumber.bool(false).asFloat(), 0.0) + func testIntValue() throws { + XCTAssertEqual(StringOrNumber.int(42).intValue, 42) + XCTAssertEqual(StringOrNumber.double(3.0).intValue, 3) // Converts + XCTAssertNil(StringOrNumber.string("hello").intValue) } - func testAsInt() throws { - XCTAssertEqual(StringOrNumber.int(42).asInt(), 42) - XCTAssertNil(StringOrNumber.float(3.14).asInt()) - XCTAssertNil(StringOrNumber.string("hello").asInt()) - XCTAssertEqual(StringOrNumber.bool(true).asInt(), 1) - XCTAssertEqual(StringOrNumber.bool(false).asInt(), 0) + func testDoubleValue() throws { + XCTAssertEqual(StringOrNumber.double(3.14).doubleValue, 3.14) + XCTAssertEqual(StringOrNumber.int(42).doubleValue, 42.0) + XCTAssertNil(StringOrNumber.string("hello").doubleValue) } - func testAsFloats() throws { - XCTAssertEqual(StringOrNumber.floats([1.1, 2.2]).asFloats(), [1.1, 2.2]) - XCTAssertEqual(StringOrNumber.ints([1, 2, 3]).asFloats(), [1.0, 2.0, 3.0]) - XCTAssertEqual(StringOrNumber.int(42).asFloats(), [42.0]) - XCTAssertNil(StringOrNumber.string("hello").asFloats()) - } - - func testAsInts() throws { - XCTAssertEqual(StringOrNumber.ints([1, 2, 3]).asInts(), [1, 2, 3]) - XCTAssertEqual(StringOrNumber.int(42).asInts(), [42]) - XCTAssertNil(StringOrNumber.floats([1.1]).asInts()) - XCTAssertNil(StringOrNumber.string("hello").asInts()) + func testFloatValue() throws { + XCTAssertEqual(StringOrNumber.double(3.14).floatValue!, 3.14, accuracy: 0.001) + XCTAssertEqual(StringOrNumber.int(42).floatValue!, 42.0, accuracy: 0.001) + XCTAssertNil(StringOrNumber.string("hello").floatValue) } // MARK: - Config Parsing Tests @@ -128,11 +90,10 @@ final class StringOrNumberTests: XCTestCase { XCTFail("Expected string type") } - XCTAssertEqual(config["factor"]?.asFloat(), 2.0) + XCTAssertEqual(config["factor"]?.floatValue, 2.0) } func testQuantizationConfig() throws { - // Typical quantization config let json = """ { "group_size": 64, @@ -145,8 +106,8 @@ final class StringOrNumberTests: XCTestCase { from: json.data(using: .utf8)! ) - XCTAssertEqual(config["group_size"]?.asInt(), 64) - XCTAssertEqual(config["bits"]?.asInt(), 4) + XCTAssertEqual(config["group_size"]?.intValue, 64) + XCTAssertEqual(config["bits"]?.intValue, 4) } // MARK: - Encoding Tests @@ -160,8 +121,8 @@ final class StringOrNumberTests: XCTestCase { let intData = try encoder.encode(StringOrNumber.int(42)) XCTAssertEqual(String(data: intData, encoding: .utf8), "42") - let boolData = try encoder.encode(StringOrNumber.bool(true)) - XCTAssertEqual(String(data: boolData, encoding: .utf8), "true") + let doubleData = try encoder.encode(StringOrNumber.double(3.14)) + XCTAssertTrue(String(data: doubleData, encoding: .utf8)?.contains("3.14") == true) } // MARK: - Equality Tests @@ -169,7 +130,22 @@ final class StringOrNumberTests: XCTestCase { func testEquality() throws { XCTAssertEqual(StringOrNumber.int(42), StringOrNumber.int(42)) XCTAssertNotEqual(StringOrNumber.int(42), StringOrNumber.int(43)) - XCTAssertNotEqual(StringOrNumber.int(42), StringOrNumber.float(42.0)) + XCTAssertNotEqual(StringOrNumber.int(42), StringOrNumber.double(42.0)) XCTAssertEqual(StringOrNumber.string("hello"), StringOrNumber.string("hello")) } + + // MARK: - Dictionary Extension Tests + + func testAsAnyDict() throws { + let dict: [String: StringOrNumber] = [ + "name": .string("test"), + "count": .int(42), + "ratio": .double(3.14), + ] + + let anyDict = dict.asAnyDict + XCTAssertEqual(anyDict["name"] as? String, "test") + XCTAssertEqual(anyDict["count"] as? Int, 42) + XCTAssertEqual(anyDict["ratio"] as? Double, 3.14) + } } diff --git a/packages/swift/Tests/NodeMLXCoreTests/SwitchLayersTests.swift b/packages/swift/Tests/NodeMLXCoreTests/SwitchLayersTests.swift new file mode 100644 index 0000000..8c864b0 --- /dev/null +++ b/packages/swift/Tests/NodeMLXCoreTests/SwitchLayersTests.swift @@ -0,0 +1,224 @@ +// Copyright © 2024 Sebastian Software GmbH. All rights reserved. +// SPDX-License-Identifier: MIT +// +// Tests for ported/SwitchLayers.swift + +import MLX +import MLXNN +import XCTest + +@testable import NodeMLXCore + +final class SwitchLayersTests: XCTestCase { + // MARK: - Helper Function Tests + + // Note: gatherSort and scatterUnsort are internal helper functions + // that are tested implicitly through the SwitchGLU tests. + // Direct testing requires very specific input formats. + + // MARK: - SwitchLinear Tests + + func testSwitchLinearInitialization() { + let layer = SwitchLinear( + inputDims: 64, + outputDims: 128, + numExperts: 4, + bias: false + ) + + XCTAssertNotNil(layer) + } + + func testSwitchLinearForward() { + let layer = SwitchLinear( + inputDims: 64, + outputDims: 128, + numExperts: 4, + bias: false + ) + + // Input: [batch*seq, topK, 1, inputDims] + let x = MLXArray.ones([8, 2, 1, 64]) + // Expert indices: [batch*seq, topK] + let indicesData: [Int32] = [0, 1, 2, 3, 0, 2, 1, 3, 0, 1, 2, 3, 0, 2, 1, 3] + let indices = MLXArray(indicesData, [8, 2]) + + let output = layer(x, indices: indices) + + // Output should be [batch*seq, topK, 1, outputDims] + XCTAssertEqual(output.shape, [8, 2, 1, 128]) + } + + func testSwitchLinearWithBias() { + let layer = SwitchLinear( + inputDims: 64, + outputDims: 128, + numExperts: 4, + bias: true + ) + + let x = MLXArray.ones([4, 2, 1, 64]) + let indicesData: [Int32] = [0, 1, 2, 3, 0, 1, 2, 3] + let indices = MLXArray(indicesData, [4, 2]) + + let output = layer(x, indices: indices) + + XCTAssertEqual(output.shape, [4, 2, 1, 128]) + } + + // MARK: - SwitchGLU Tests + + func testSwitchGLUInitialization() { + let glu = SwitchGLU( + inputDims: 64, + hiddenDims: 256, + numExperts: 4, + bias: false + ) + + XCTAssertNotNil(glu) + } + + func testSwitchGLUForward() { + let glu = SwitchGLU( + inputDims: 64, + hiddenDims: 256, + numExperts: 4, + bias: false + ) + + // Input: [batch, seq, hidden] + let x = MLXArray.ones([2, 8, 64]) + // Expert indices: [batch, seq, topK] + let indices = MLXArray([Int32](repeating: 0, count: 16) + [Int32](repeating: 1, count: 16)).reshaped([2, 8, 2]) + + let output = glu(x, indices: indices) + + // Output should be [batch, seq, topK, hidden] + XCTAssertEqual(output.dim(0), 2) + XCTAssertEqual(output.dim(1), 8) + XCTAssertEqual(output.dim(-1), 64) + } + + // MARK: - SwiGLU Activation Tests + + func testSwiGLU() { + let x = MLXArray.ones([4, 64]) + let gate = MLXArray.ones([4, 64]) + + let output = swiGLU(x, gate: gate) + + XCTAssertEqual(output.shape, x.shape) + } + + // MARK: - SwiGLUSwitchGLU Tests (GPT-OSS) + + func testSwiGLUSwitchGLUInitialization() { + let glu = SwiGLUSwitchGLU( + inputDims: 64, + hiddenDims: 256, + numExperts: 4, + bias: true + ) + + XCTAssertNotNil(glu) + } + + func testSwiGLUSwitchGLUForward() { + let glu = SwiGLUSwitchGLU( + inputDims: 64, + hiddenDims: 256, + numExperts: 4, + bias: true + ) + + let x = MLXArray.ones([2, 8, 64]) + let indices = MLXArray([Int32](repeating: 0, count: 16) + [Int32](repeating: 1, count: 16)).reshaped([2, 8, 2]) + + let output = glu(x, indices: indices) + + XCTAssertEqual(output.dim(0), 2) + XCTAssertEqual(output.dim(1), 8) + XCTAssertEqual(output.dim(-1), 64) + } + + // MARK: - SwitchMLP Tests + + func testSwitchMLPInitialization() { + let mlp = SwitchMLP( + inputDims: 64, + hiddenDims: 256, + numExperts: 4 + ) + + XCTAssertNotNil(mlp) + } + + func testSwitchMLPForward() { + let mlp = SwitchMLP( + inputDims: 64, + hiddenDims: 256, + numExperts: 4, + activation: gelu + ) + + let x = MLXArray.ones([2, 8, 64]) + let indices = MLXArray([Int32](repeating: 0, count: 16) + [Int32](repeating: 1, count: 16)).reshaped([2, 8, 2]) + + let output = mlp(x, indices: indices) + + XCTAssertEqual(output.dim(0), 2) + XCTAssertEqual(output.dim(1), 8) + XCTAssertEqual(output.dim(-1), 64) + } + + // MARK: - MoE Tensor Conversion Tests + + func testConvertMoePackedTensors() { + // Create mock packed tensors + let blocks = MLXArray.ones([4, 64, 256]) + let scales = MLXArray.ones([4, 64, 4]) + + let result = convertMoePackedTensors(blocks: blocks, scales: scales) + + // Result should maintain expert dimension + XCTAssertEqual(result.dim(0), 4) + } + + // MARK: - Edge Cases + + func testSwitchLinearSingleExpert() { + let layer = SwitchLinear( + inputDims: 64, + outputDims: 128, + numExperts: 1, + bias: false + ) + + let x = MLXArray.ones([4, 1, 1, 64]) + let indicesData: [Int32] = [0, 0, 0, 0] + let indices = MLXArray(indicesData, [4, 1]) + + let output = layer(x, indices: indices) + + XCTAssertEqual(output.shape, [4, 1, 1, 128]) + } + + func testSwitchLinearManyExperts() { + let layer = SwitchLinear( + inputDims: 64, + outputDims: 128, + numExperts: 16, + bias: false + ) + + let x = MLXArray.ones([8, 4, 1, 64]) + // Expert indices between 0-15 + let indicesData: [Int32] = (0 ..< 32).map { Int32($0 % 16) } + let indices = MLXArray(indicesData, [8, 4]) + + let output = layer(x, indices: indices) + + XCTAssertEqual(output.shape, [8, 4, 1, 128]) + } +} diff --git a/packages/swift/Tests/NodeMLXCoreTests/TokenizerTests.swift b/packages/swift/Tests/NodeMLXCoreTests/TokenizerTests.swift deleted file mode 100644 index c80fa82..0000000 --- a/packages/swift/Tests/NodeMLXCoreTests/TokenizerTests.swift +++ /dev/null @@ -1,94 +0,0 @@ -import Hub -@testable import NodeMLXCore -import Tokenizers -import XCTest - -final class TokenizerTests: XCTestCase { - // MARK: - Basic Tokenizer Tests - - func testHFTokenizerFromHub() async throws { - // Test mit Qwen - modernes Modell mit korrektem Config-Format - let tokenizer = try await HFTokenizer(modelId: "Qwen/Qwen2.5-0.5B-Instruct") - - let text = "Hello, world!" - let tokens = tokenizer.encode(text) - - XCTAssertFalse(tokens.isEmpty, "Tokens should not be empty") - print("✓ Encoded '\(text)' to \(tokens.count) tokens: \(tokens)") - - let decoded = tokenizer.decode(tokens) - print("✓ Decoded back to: '\(decoded)'") - - XCTAssertTrue(decoded.contains("Hello")) - } - - func testHFHubCachePath() { - // Test cache path generation - let path = HFHubCache.modelPath(for: "mlx-community/Llama-3.2-1B-Instruct-4bit") - XCTAssertTrue(path.path.contains("mlx-community--Llama-3.2-1B-Instruct-4bit")) - print("✓ Cache path: \(path.path)") - } - - func testTokenizerRoundtrip() async throws { - // Test mit Qwen tokenizer (nicht gated!) - let tokenizer = try await HFTokenizer(modelId: "Qwen/Qwen2.5-0.5B-Instruct") - - let texts = [ - "Hello!", - "The quick brown fox jumps over the lazy dog.", - "1 + 1 = 2", - ] - - for text in texts { - let tokens = tokenizer.encode(text) - let decoded = tokenizer.decode(tokens) - print("✓ '\(text)' → \(tokens.count) tokens → '\(decoded)'") - - XCTAssertFalse(tokens.isEmpty) - } - } - - // MARK: - Chat Template Tests (important for LLM inference) - - func testChatTemplateAvailability() async throws { - // AutoTokenizer from swift-transformers should support chat templates - let hub = HubApi() - let repo = Hub.Repo(id: "Qwen/Qwen2.5-0.5B-Instruct") - let modelDir = try await hub.snapshot(from: repo, matching: ["tokenizer*", "vocab*", "merges*"]) - - let tokenizer = try await AutoTokenizer.from(modelFolder: modelDir) - - // Check if chat template is available - let messages: [[String: String]] = [ - ["role": "user", "content": "Hello!"], - ] - - // Try to apply chat template - do { - let result = try tokenizer.applyChatTemplate(messages: messages) - print("✓ Chat template applied, got \(result.count) tokens") - XCTAssertFalse(result.isEmpty) - } catch { - print("⚠ Chat template not available: \(error)") - // Not all tokenizers have chat templates, so this isn't necessarily a failure - } - } - - // MARK: - Special Tokens Tests - - func testSpecialTokens() async throws { - let tokenizer = try await HFTokenizer(modelId: "Qwen/Qwen2.5-0.5B-Instruct") - - print("Special tokens:") - print(" BOS: \(tokenizer.bosTokenId ?? -1)") - print(" EOS: \(tokenizer.eosTokenId ?? -1)") - print(" PAD: \(tokenizer.padTokenId ?? -1)") - - // At least one special token should be defined - let hasSpecialTokens = tokenizer.bosTokenId != nil || - tokenizer.eosTokenId != nil || - tokenizer.padTokenId != nil - - print("✓ Has special tokens: \(hasSpecialTokens)") - } -} diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index a6017b7..95ab4d8 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -155,6 +155,9 @@ importers: tsup: specifier: ^8.5.1 version: 8.5.1(jiti@2.6.1)(postcss@8.5.6)(tsx@4.21.0)(typescript@5.9.3)(yaml@2.8.2) + tsx: + specifier: ^4.21.0 + version: 4.21.0 typescript: specifier: ^5.9.3 version: 5.9.3