Instructions to use AiArtLab/sdxs-1b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use AiArtLab/sdxs-1b with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("AiArtLab/sdxs-1b", dtype=torch.bfloat16, device_map="cuda") prompt = "sdxs-1b" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
Download migrate_tapered.py from AiArtLab/sdxs-1b: direct link, hf CLI and curl.
- Browser
- Download file 3.98 kB
-
https://huggingface.co/AiArtLab/sdxs-1b/resolve/main/migrate_tapered.py
- Command line
-
hf download hf://AiArtLab/sdxs-1b/migrate_tapered.py
-
curl -L -o migrate_tapered.py https://huggingface.co/AiArtLab/sdxs-1b/resolve/main/migrate_tapered.py
3.98 kB
| import os, json, torch | |
| from safetensors.torch import load_file, save_file | |
| from diffusers import UNet2DConditionModel | |
| SRC_DIR = "unet" # trained uniform model (layers=2, trans=[3,3,3,3]) | |
| DST_DIR = "unet2" # new tapered model | |
| TAPERED_CFG = { | |
| "block_out_channels": [320, 640, 1280, 1536], | |
| "layers_per_block": [ 4, 3, 2, 2], | |
| "transformer_layers_per_block": [ 4, 3, 2, 1], | |
| "attention_head_dim": [ 10, 10, 10, 12], | |
| } | |
| os.makedirs(DST_DIR, exist_ok=True) | |
| # 1. Load source config, apply tapered overrides | |
| with open(os.path.join(SRC_DIR, "config.json"), "r", encoding="utf-8") as f: | |
| cfg = json.load(f) | |
| for k in ["_class_name", "_diffusers_version", "_name_or_path"]: | |
| cfg.pop(k, None) | |
| cfg.update(TAPERED_CFG) | |
| with open(os.path.join(DST_DIR, "config.json"), "w", encoding="utf-8") as f: | |
| json.dump({ | |
| "_class_name": "UNet2DConditionModel", | |
| "_diffusers_version": "0.37.1", | |
| "_name_or_path": "unet2", | |
| **cfg, | |
| }, f, indent=2) | |
| # 2. Load old weights | |
| print("Loading old weights...") | |
| old_state = load_file(os.path.join(SRC_DIR, "diffusion_pytorch_model.safetensors"), device="cpu") | |
| # 3. Create new model | |
| print("Creating tapered model...") | |
| new_model = UNet2DConditionModel(**cfg) | |
| new_state = new_model.state_dict() | |
| old_params = sum(p.numel() for p in old_state.values()) / 1e9 | |
| new_params = sum(p.numel() for p in new_state.values()) / 1e9 | |
| print(f"Old: {old_params:.3f}B → New: {new_params:.3f}B") | |
| # 4. Copy exact matches | |
| copied, warm = 0, [] | |
| for key, old_t in old_state.items(): | |
| if key in new_state and old_t.shape == new_state[key].shape: | |
| new_state[key] = old_t.clone() | |
| copied += 1 | |
| # 5. Warm-start NEW layers from closest existing layer | |
| # Strategy: for each block, if a resnet.X or transformer.X is new (exists in new but not in old), | |
| # try to copy from X-1. If X-1 also doesn't exist, try X-2, etc. | |
| def find_source(old_state, key, new_state): | |
| """Try to find a source for a new key by walking backwards through indices.""" | |
| # Extract prefix and index: e.g. "down_blocks.0.resnets.2." → ("down_blocks.0.resnets.", 2) | |
| import re | |
| parts = key.rsplit(".", 2) | |
| if len(parts) < 3: | |
| return None | |
| prefix = parts[0] + "." # e.g. "down_blocks.0.resnets." | |
| # Find the numeric index before the last segment | |
| m = re.match(r"(.*\.)(\d+)\.(.+)", key) | |
| if not m: | |
| return None | |
| stem, idx_str, suffix = m.group(1), m.group(2), m.group(3) | |
| idx = int(idx_str) | |
| for src_idx in range(idx - 1, -1, -1): | |
| src_key = f"{stem}{src_idx}.{suffix}" | |
| if src_key in old_state and old_state[src_key].shape == new_state[key].shape: | |
| return src_key | |
| return None | |
| warm_started = [] | |
| for key in new_state: | |
| if key not in old_state or old_state[key].shape != new_state[key].shape: | |
| src_key = find_source(old_state, key, new_state) | |
| if src_key: | |
| new_state[key] = old_state[src_key].clone() | |
| warm_started.append((key, src_key)) | |
| print(f"Direct copies: {copied}") | |
| print(f"Warm-started: {len(warm_started)} keys") | |
| for k, src in warm_started[:10]: | |
| print(f" {k} ← {src}") | |
| if len(warm_started) > 10: | |
| print(f" ... and {len(warm_started) - 10} more") | |
| # 6. Report stats | |
| total_copied = sum(old_state[k].numel() for k in old_state if k in new_state and old_state[k].shape == new_state[k].shape) | |
| total_ws = sum(new_state[k].numel() for k, _ in warm_started) | |
| total_new = sum(v.numel() for v in new_state.values()) | |
| print(f"\nCopied: {total_copied/1e9:.3f}B ({total_copied/total_new*100:.1f}%)") | |
| print(f"Warm-started: {total_ws/1e9:.3f}B ({total_ws/total_new*100:.1f}%)") | |
| print(f"Random init: {(total_new - total_copied - total_ws)/1e9:.3f}B ({(total_new - total_copied - total_ws)/total_new*100:.1f}%)") | |
| # 7. Save | |
| print("\nSaving...") | |
| save_file(new_state, os.path.join(DST_DIR, "diffusion_pytorch_model.safetensors")) | |
| print(f"Done → {DST_DIR}/") | |