Metadata-Version: 2.4
Name: zadapter
Version: 0.2.0
Summary: JAX/Flax-native bottleneck FFN adapter with Activation Function Annealing (AFA), tensor-parallel sharding, and zenbit NF4 quantization integration. TPU-optimized PEFT alternative to LoRA.
Author-email: Riki <alkenzyhaikal@gmail.com>
License: MIT
Project-URL: Homepage, https://github.com/RikZD/ZAdapter
Project-URL: Repository, https://github.com/RikZD/ZAdapter
Project-URL: Documentation, https://github.com/RikZD/ZAdapter#readme
Project-URL: Changelog, https://github.com/RikZD/ZAdapter/blob/main/CHANGELOG.md
Keywords: jax,flax,pallas,tpu,xla,fine-tuning,peft,adapter,lora,tensor-parallel,nf4,quantization
Classifier: Development Status :: 4 - Beta
Classifier: Intended Audience :: Developers
Classifier: Intended Audience :: Science/Research
Classifier: License :: OSI Approved :: MIT License
Classifier: Programming Language :: Python :: 3
Classifier: Programming Language :: Python :: 3.9
Classifier: Programming Language :: Python :: 3.10
Classifier: Programming Language :: Python :: 3.11
Classifier: Programming Language :: Python :: 3.12
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Requires-Python: >=3.9
Description-Content-Type: text/markdown
License-File: LICENSE
Requires-Dist: jax>=0.4.20
Requires-Dist: flax>=0.7.0
Provides-Extra: jax
Provides-Extra: pallas
Requires-Dist: zenbit[pallas-flax]>=0.2.0; extra == "pallas"
Provides-Extra: pytorch
Requires-Dist: torch>=2.0; extra == "pytorch"
Provides-Extra: all
Requires-Dist: zenbit[pallas-flax]>=0.2.0; extra == "all"
Requires-Dist: torch>=2.0; extra == "all"
Dynamic: license-file

# ZAdapter

TPU-native bottleneck FFN adapter for parameter-efficient fine-tuning.

Unlike LoRA (parallel low-rank matrices merged into existing weights),
ZAdapter injects a small bottleneck feed-forward block **serially** after
attention and MLP in each transformer layer. This keeps the computation
graph static — no merge/unmerge step, no dynamic branching — which plays
well with XLA compilation on TPU.

```
input -> down_proj [d_model -> r] -> activation -> up_proj [r -> d_model] -> output
output = input + adapter(input)
```

## Install

```bash
# JAX/Flax native (recommended for TPU)
pip install zadapter

# With zenbit NF4/Pallas integration
pip install zadapter[pallas]

# With PyTorch backward compatibility
pip install zadapter[pytorch]

# Everything
pip install zadapter[all]
```

## Quick Start (JAX/Flax)

```python
import jax
import jax.numpy as jnp
import flax.linen as nn
from zadapter import ZAdapter, AFAScheduleConfig

# Inside your Flax module
class MyBlock(nn.Module):
    d_model: int = 4096
    r: int = 64

    @nn.compact
    def __call__(self, x, step: int = 0):
        # ... attention, MLP, etc ...

        # Inject ZAdapter with Activation Function Annealing (AFA)
        schedule = AFAScheduleConfig(total_steps=10000, anneal_fraction=0.3)
        beta = schedule.beta(step)
        x = ZAdapter(d_model=self.d_model, r=self.r)(x, beta=beta)

        return x
```

## Activation Function Annealing (AFA)

ZAdapter v0.2 introduces **AFA** — adapted from *Li et al., "AFA-LoRA: Enabling Non-Linear Adaptations in LoRA with Activation Function Annealing" (2026)*.

- **Training**: Uses non-linear activation (GELU/ReLU/SiLU) for expressiveness
- **Annealing**: Smoothly transitions to identity function over first ~30% of training
- **Post-training**: Adapter collapses to a pure linear map → **mergeable** into base weight!

```python
from zadapter import merge_into_weight

# After training (beta fully annealed to 0)
merged_weight = merge_into_weight(base_weight, adapter_params)
```

**Shard-aware merge**: Works on already-sharded weights (tensor-parallel) — no need to gather 900B+ models onto a single device.

## Tensor Parallel (Large Models)

For models too large for a single TPU core (~30B+ params even with NF4):

```python
from zadapter import setup_mesh, ColumnParallelDense, RowParallelDense, TPConfig

mesh = setup_mesh(TPConfig(num_shards=8))

class ShardedBlock(nn.Module):
    @nn.compact
    def __call__(self, x, beta=0.0):
        x = ColumnParallelDense(features=4*d_model, mesh=mesh)(x)
        x = ZAdapter(d_model=4*d_model, r=64)(x, beta=beta)
        x = RowParallelDense(features=d_model, mesh=mesh)(x)
        return x
```

ZAdapter params are **replicated** (not sharded) across the mesh — negligible memory cost, zero cross-device communication overhead for something that small.

## Integration with zenbit (NF4 Quantization)

```python
from zenbit.pallas_nf4.flax_layer import NF4DenseFused
from zadapter import ZAdapter

# NF4-quantized frozen base + trainable ZAdapter on top
h = NF4DenseFused(features=d_model)(x, quantized_weight)
h = ZAdapter(d_model=d_model, r=64)(h, beta=1.0)
```

See `tests/test_zadapter_zenbit_integration.py` for full integration tests.

## Why ZAdapter over LoRA

| | LoRA | ZAdapter v0.1 (PyTorch) | ZAdapter v0.2 (JAX) |
|---|---|---|---|
| Injection | Parallel to weight | Serial FFN block | Serial FFN block |
| Mergeable | Yes (after training) | **No** | **Yes** (via AFA) |
| Graph | Dynamic (merge/unmerge) | Static, XLA-friendly | Static, XLA-friendly |
| Trainable params | ~0.1-1% | ~0.5-3% | ~0.5-3% |
| Backend | PyTorch | PyTorch | **JAX/Flax/Pallas** |
| Tensor Parallel | Manual | — | **Native (Megatron-style)** |
| Quantization | bitsandbytes | — | **zenbit NF4 (Pallas kernel)** |

## Legacy PyTorch API (v0.1)

The PyTorch version is still available for backward compatibility:

```python
from zadapter.pytorch import inject_adapter, get_trainable_params
```

Install with: `pip install zadapter[pytorch]`

## Companion Library

For TPU sharding, data-parallel training loops, and NF4/int8 quantization:
[zenbit](https://pypi.org/project/zenbit/)

## License

MIT
