diff --git a/Llama3_1_(8B)_Alpaca-ASFT.ipynb b/Llama3_1_(8B)_Alpaca-ASFT.ipynb new file mode 100644 index 0000000000..987db4d3c0 --- /dev/null +++ b/Llama3_1_(8B)_Alpaca-ASFT.ipynb @@ -0,0 +1,4826 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": { + "id": "x_wPZgziQKXy" + }, + "source": [ + "To run this, press \"*Runtime*\" and press \"*Run all*\" on a **free** Tesla T4 Google Colab instance!\n", + "
\n", + "\n", + "\n", + " Join Discord if you need help + ⭐ Star us on Github ⭐\n", + "
\n", + "\n", + "This notebook is an **ASFT / ASFT+ demo** (Anchored Supervised Fine-Tuning).\n", + "\n", + "Credits:\n", + "- ASFT paper & reference implementation: https://github.com/zhuchichi56/ASFT\n", + "- ASFT+ (this optimized Unsloth integration + extra speed/perf optimizations): Hasan Can Solakoğlu X/Twitter @HCSolakoglu\n", + "\n", + "To install Unsloth your local device, follow [our guide](https://docs.unsloth.ai/get-started/install-and-update). This notebook is licensed [LGPL-3.0](https://github.com/unslothai/notebooks?tab=LGPL-3.0-1-ov-file#readme).\n", + "\n", + "You will learn how to do [data prep](#Data), how to [train](#Train), how to [run the model](#Inference), & [how to save it](#Save)\n" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "t-ahwuyvQKXz" + }, + "source": [ + "### News" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "TiJUkQ5MQKX0" + }, + "source": [ + "\n", + "Introducing FP8 precision training for faster RL inference. [Read Blog](https://docs.unsloth.ai/new/fp8-reinforcement-learning).\n", + "\n", + "Unsloth's [Docker image](https://hub.docker.com/r/unsloth/unsloth) is here! Start training with no setup & environment issues. [Read our Guide](https://docs.unsloth.ai/new/how-to-train-llms-with-unsloth-and-docker).\n", + "\n", + "[gpt-oss RL](https://docs.unsloth.ai/new/gpt-oss-reinforcement-learning) is now supported with the fastest inference & lowest VRAM. Try our [new notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/gpt-oss-(20B)-GRPO.ipynb) which creates kernels!\n", + "\n", + "Introducing [Vision](https://docs.unsloth.ai/new/vision-reinforcement-learning-vlm-rl) and [Standby](https://docs.unsloth.ai/basics/memory-efficient-rl) for RL! Train Qwen, Gemma etc. VLMs with GSPO - even faster with less VRAM.\n", + "\n", + "Visit our docs for all our [model uploads](https://docs.unsloth.ai/get-started/all-our-models) and [notebooks](https://docs.unsloth.ai/get-started/unsloth-notebooks).\n" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "vXSL0oj6QKX0" + }, + "source": [ + "### Installation" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "krAQhz2yQKX0" + }, + "outputs": [], + "source": [ + "%%capture\n", + "import os, re\n", + "\n", + "IN_COLAB = \"COLAB_\" in \"\".join(os.environ.keys())\n", + "\n", + "if not IN_COLAB:\n", + " # ASFT demo: if you're running this notebook from the Unsloth repo/branch,\n", + " # an editable install is the most reliable way to ensure ASFTTrainer is present.\n", + " if os.path.exists(\"pyproject.toml\"):\n", + " %pip install -e \".[cu126-torch290]\"\n", + " else:\n", + " %pip install -U unsloth\n", + "else:\n", + " # Do this only in Colab notebooks! Otherwise use pip install unsloth / pip install -e .\n", + " import torch; v = re.match(r\"[0-9]{1,}\\.[0-9]{1,}\", str(torch.__version__)).group(0)\n", + " xformers = \"xformers==\" + (\"0.0.33.post1\" if v==\"2.9\" else \"0.0.32.post2\" if v==\"2.8\" else \"0.0.29.post3\")\n", + " %pip install --no-deps bitsandbytes accelerate {xformers} peft trl triton cut_cross_entropy unsloth_zoo\n", + " %pip install sentencepiece protobuf \"datasets==4.3.0\" \"huggingface_hub>=0.34.0\" hf_transfer\n", + " %pip install --no-deps unsloth\n", + "\n", + "%pip install transformers==4.56.2\n", + "%pip install --no-deps trl==0.22.2" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# In Colab, clone the ASFT branch of Unsloth for ASFTTrainer\n", + "!git clone -b asft-plus https://github.com/hcsolakoglu/unsloth.git\n", + "%cd unsloth" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "!pip install -e \".[cu126-torch290]\"" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "S68-v0avQKX1" + }, + "source": [ + "### Unsloth" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 300, + "referenced_widgets": [ + "6e3c281f112b4a86af7a3ef95933d221", + "92395f250a154006923aaf9ea0a9c30b", + "f84bfc5390054ec687c157c4d68199a6", + "4228734651ca45e19fc7bda79817f9b3", + "4613edbbec6846edb5b1677c25d542b6", + "d032fe2ba5d647d99026fdade758c0cd", + "e9971d220fe24552a1e9aa299765cfb9", + "2be29a4553ad4dfea8a9bc620c81a3ae", + "94cbb87829d1486899e2ff6325c2ecdf", + "f878c2e00bc240c7b0333cce950080e1", + "6b908368de51428585552dfef6a83088", + "ac5eacaaee8346c080e54ea7a52648a4", + "7981edf408d54d41bbeac42da7492c6b", + "2b848e5a85bc42bc87945fd9ed5db038", + "e88c33f37d6849e0b1a6b41254104cb9", + "33843107b93647b28985bfc37ea781ca", + "76e5410e286a4a5abd6c213a38aa38bb", + "ae7e90a811f94e75997d6a9ed1be8596", + "bb9f3379310d4b04be694996f3137b28", + "75082ba15db445df907f5612976590ae", + "89934c4f26834f15b9889ec36fee3b65", + "e887160635cb4803b9f33845df615ec6", + "1c7bc5fdb7dd4c39af8d4c2c504ec3ed", + "843a27e619534ea8914f9d36386c364b", + "8f5adc70fbf248f2811527f620553be5", + "4ffb4b2f015046fb94c1115ed0397a20", + "5a232ed040f94633a2a374031284c1f6", + "2006be31c09349738e221295bb84939f", + "07055fc12b0841aaa5317f8252b5d347", + "1eb90e686e214122ae763b1b79ae321d", + "7a935956348e47c68fbdf05ddf4752f3", + "8653acb618ad4e76bbf1daa00ea71238", + "e3c3bd9c4c124b0a8c88c83c1fc747d3", + "1bd75ddaf57c4438a4e2c3070b9cef65", + "a3a3ef6d6337403cabea8b23f7c3021b", + "c2ea0a3f01f34ffa8c94ab9b5098e9da", + "68ea1d7cb8274a639b3fb5326f4218c3", + "39fef7b257614a0595f39355fa226b69", + "d125995cc0934239a01ba01b78529f21", + "634ae4c6cfe04673b1cdc9c9cac4cbf9", + "d7f92e8332374313bee87ccd427446a4", + "36799fbcd90d43128620ff98225a825d", + "5310346dd579424fa676b8e8e64790e7", + "0c1835f404db4846bb13b5da8d8f4447", + "29c5b713f07043dda51820523e5c8ff3", + "d7375f0f048841b29a20601c122666e8", + "f433ced9bfcd4a57ba691d3c1caeed08", + "da8ffc70820a48f5a12c6d4b5967015b", + "1a6db9aea6a64ae3aaef51d6265b35b2", + "331f516c7a76456d801bc2a2feb228aa", + "9be9074028da42d39d044a78393a861f", + "51cd9026b1664819a67712996ca97bd5", + "ea55293415ca48a4be97c2e1e4769122", + "8ea52b105a7e44978caca33c0e7e815b", + "3662f1445ef34a50b462e601ed31bb69" + ] + }, + "id": "QmUBVEnvCDJv", + "outputId": "0a47b925-663d-4543-9c61-994a6302f3c5" + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "🦥 Unsloth: Will patch your computer to enable 2x faster free finetuning.\n", + "==((====))== Unsloth 2024.8: Fast Llama patching. Transformers = 4.44.2.\n", + " \\\\ /| GPU: Tesla T4. Max memory: 14.748 GB. Platform = Linux.\n", + "O^O/ \\_/ \\ Pytorch: 2.4.0+cu121. CUDA = 7.5. CUDA Toolkit = 12.1.\n", + "\\ / Bfloat16 = FALSE. FA [Xformers = 0.0.27.post2. FA2 = False]\n", + " \"-____-\" Free Apache license: http://github.com/unslothai/unsloth\n", + "Unsloth: Fast downloading is enabled - ignore downloading bars which are red colored!\n" + ] + }, + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "6e3c281f112b4a86af7a3ef95933d221", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "model.safetensors: 0%| | 0.00/5.70G [00:00 0 ! Suggested 8, 16, 32, 64, 128\n", + " target_modules = [\"q_proj\", \"k_proj\", \"v_proj\", \"o_proj\",\n", + " \"gate_proj\", \"up_proj\", \"down_proj\",],\n", + " lora_alpha = 16,\n", + " lora_dropout = 0, # Supports any, but = 0 is optimized\n", + " bias = \"none\", # Supports any, but = \"none\" is optimized\n", + " # [NEW] \"unsloth\" uses 30% less VRAM, fits 2x larger batch sizes!\n", + " use_gradient_checkpointing = \"unsloth\", # True or \"unsloth\" for very long context\n", + " random_state = 3407,\n", + " use_rslora = False, # We support rank stabilized LoRA\n", + " loftq_config = None, # And LoftQ\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "vITh0KVJ10qX" + }, + "source": [ + "\n", + "### Data Prep\n", + "We now use the Alpaca dataset from [yahma](https://huggingface.co/datasets/yahma/alpaca-cleaned), which is a filtered version of 52K of the original [Alpaca dataset](https://crfm.stanford.edu/2023/03/13/alpaca.html). You can replace this code section with your own data prep.\n", + "\n", + "**[NOTE]** To train only on completions (ignoring the user's input) read TRL's docs [here](https://huggingface.co/docs/trl/sft_trainer#train-on-completions-only).\n", + "\n", + "**[NOTE]** Remember to add the **EOS_TOKEN** to the tokenized output!! Otherwise you'll get infinite generations!\n", + "\n", + "If you want to use the `llama-3` template for ShareGPT datasets, try our conversational [notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3_(8B)-Alpaca.ipynb)\n", + "\n", + "For text completions like novel writing, try this [notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Mistral_(7B)-Text_Completion.ipynb)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 145, + "referenced_widgets": [ + "5e8825fb770b41529f2129113cebc4a9", + "0a9dc233674e4096b7a988a5e4ebaf84", + "374fa9beda4042e1bf9a9b13de6e6674", + "e6533d3c91fd4359bc84ffd8e59af5a3", + "f9b01aebcbdc48a585b7942b0ee60a2d", + "d4b3770433bc41818372b7aed243fb31", + "19bcefcc1d874840ae9a9ca983e474b6", + "4011ce9370d74fad857ec8e1e99d314f", + "41f5fed060ad4c8d87b24602b720ef04", + "88fcd51819b5483c9ab22df7ef89ab64", + "e6ffac074f1b476ba2ade11b37732af3", + "98a6716e7438429ea322adb3e3264f91", + "68c686291b50430faeef0de7840e2c4b", + "953625aa1e824f8a8d203197b316b302", + "f899a815142542219bde22ff792fb60c", + "51ca174d26e94b5cb1e895aa3c770655", + "f4519637bb43400a80ce83505101e8a5", + "805676b197c94f5aa45956daa354640b", + "ce9fcc5eff1f460d80b703a4ca32dad1", + "3fef797403d14440afe599a3bf06b626", + "4e7cb8e988114ed4b6fe09ff9f682dff", + "a14bc1c2130842568a5fde6698731e5f", + "85cc6f24cba54563acb5598f54fed7b9", + "80a72037771e4da9be989eefabbc8e76", + "ba68b274c50b44ec9e02642378d271a6", + "f84d2fe4f1c24a34948755abf1f32b7f", + "86511967834f4484a5ec4af387b7d7a9", + "c97c40c2bf2a41a8ae1c75e0a9c8ebff", + "484e507f14424f2b9173595b985f4101", + "68a32b398e1c490393e01befdc260785", + "04cc963133d242779572d2e847fa3d65", + "94730f13e92a4c9aac35c2cfb21fc48c", + "620c0de28ec74f71a021a2be96dccf3a", + "6e1aff64771c402ab070f650562fa4c9", + "0078f897f2174217a307d95d4f9bd775", + "ecdeaab4f8c94d6dade63bb06857c969", + "735b85f0a0e9411cac4d704a504fcfc1", + "8e992e60416145a8b6eed744287ca0fb", + "8c195b5809604905b5e404baa30e8449", + "2247efac4283489bbd228330344388ab", + "e93d063faf984cc4aa51462418d9b57e", + "9d45b9a5de3e4cba9ac35ad2cb187f51", + "e99423a1ed3f4f72886b39368468b7c1", + "5f174718e5974a7cab024d113f662513" + ] + }, + "id": "LjY75GoYUCB8", + "outputId": "80d6c3b9-28c2-4ebf-9c57-6a0b77ce82b1" + }, + "outputs": [ + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "5e8825fb770b41529f2129113cebc4a9", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Downloading readme: 0%| | 0.00/11.6k [00:00\n", + "### Train the model (ASFT / ASFT+ demo)\n", + "This demo notebook uses **Unsloth `ASFTTrainer`** (Anchored Supervised Fine-Tuning) instead of the standard `SFTTrainer`.\n", + "\n", + "ASFT in a nutshell:\n", + "- Uses **DFT weights** (based on token probabilities / confidence) to reweight token-level CE loss.\n", + "- Adds lightweight **KL anchoring** to stay close to a reference distribution (stability).\n", + "- Supports **streaming** to chunk the reference forward pass and reduce peak VRAM.\n", + "\n", + "**ASFT+** in this repo refers to the same ASFT objective with extra engineering work (performance + VRAM optimizations) on top.\n", + "\n", + "Note: `max_steps` is kept small for a quick demo. For a full run, set `max_steps=None` and use `num_train_epochs=1` (or similar)." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "#### Recommended ASFT defaults (optimized)\n", + "- `asft_mode=\"asft\"` + `kl_weight=0.05`: a solid starting point (stable, not overly restrictive).\n", + "- `reference_policy=\"disable_adapter\"`: avoids keeping a separate frozen reference copy in most PEFT setups.\n", + "- `ASFTStreamingConfig(mode=\"auto\")`: tries seq_kv_cache first, falls back to batch micro to reduce peak VRAM.\n", + "\n", + "For quick comparisons, try: `asft_mode=\"sft\"` or `asft_mode=\"dft\"` (KL off)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/", + "height": 67, + "referenced_widgets": [ + "3719bf6f9c6a4c6fbef93c5328c11a07", + "03f492b4b56f4d8e80e9395a65058b1b", + "39d9ef9fb35f47119f319f48eb222070", + "3d7cfb33ceaf417e851ac4393c65148b", + "ece66fa2f128456fa2a82b8a28d1211c", + "9695a640b0ff4e91af495bb59548e4b6", + "d4bd5559d4134d64a943d57972c6ef39", + "fbae6e599d1644f39e5d86efa0f9f997", + "00d425bca350451da6400f9f05c4a659", + "6a27d9ad4f064586a87636b10455d15b", + "77f4367616964a01a8c42416f5f4c147" + ] + }, + "id": "95_Nn-89DhsL", + "outputId": "29798478-b975-42d3-b32b-020a805cac35" + }, + "outputs": [ + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "3719bf6f9c6a4c6fbef93c5328c11a07", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Map (num_proc=2): 0%| | 0/51760 [00:00\n", + " \n", + " \n", + " [60/60 07:28, Epoch 0/1]\n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + " \n", + "
StepTraining Loss
11.817600
22.304200
31.689300
41.938200
51.656900
61.621900
71.187100
81.264200
91.101200
101.189500
110.930800
120.959400
130.929400
141.048700
150.892800
160.901400
171.009100
181.256100
191.016500
200.882600
210.940500
221.018500
230.897200
240.991900
251.072000
261.022900
271.044900
280.877800
290.843800
300.887500
310.853400
320.866000
330.983200
340.852200
350.961200
360.856700
370.872300
380.751100
391.081400
401.174400
410.893400
420.977500
430.957100
440.908100
450.915000
460.973400
470.870900
481.196500
490.907500
501.031300
511.015900
520.907900
530.977000
541.154300
550.778000
561.013300
570.886800
580.827500
590.852300
600.896600

" + ], + "text/plain": [ + "" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "trainer_stats = trainer.train()" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "cellView": "form", + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "pCqnaKmlO1U9", + "outputId": "edf33a96-b12c-4bba-9771-59e18aee707c" + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "462.7198 seconds used for training.\n", + "7.71 minutes used for training.\n", + "Peak reserved memory = 7.922 GB.\n", + "Peak reserved memory for training = 1.938 GB.\n", + "Peak reserved memory % of max memory = 53.716 %.\n", + "Peak reserved memory for training % of max memory = 13.141 %.\n" + ] + } + ], + "source": [ + "# @title Show final memory and time stats\n", + "used_memory = round(torch.cuda.max_memory_reserved() / 1024 / 1024 / 1024, 3)\n", + "used_memory_for_lora = round(used_memory - start_gpu_memory, 3)\n", + "used_percentage = round(used_memory / max_memory * 100, 3)\n", + "lora_percentage = round(used_memory_for_lora / max_memory * 100, 3)\n", + "print(f\"{trainer_stats.metrics['train_runtime']} seconds used for training.\")\n", + "print(\n", + " f\"{round(trainer_stats.metrics['train_runtime']/60, 2)} minutes used for training.\"\n", + ")\n", + "print(f\"Peak reserved memory = {used_memory} GB.\")\n", + "print(f\"Peak reserved memory for training = {used_memory_for_lora} GB.\")\n", + "print(f\"Peak reserved memory % of max memory = {used_percentage} %.\")\n", + "print(f\"Peak reserved memory for training % of max memory = {lora_percentage} %.\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "ekOmTR1hSNcr" + }, + "source": [ + "\n", + "### Inference\n", + "Let's run the model! You can change the instruction and input - leave the output blank!\n", + "\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "kR3gIAX-SM2q", + "outputId": "087c5c13-e946-4c35-e4f2-e07a88f9ac32" + }, + "outputs": [ + { + "data": { + "text/plain": [ + "['<|begin_of_text|>Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.\\n\\n### Instruction:\\nContinue the fibonnaci sequence.\\n\\n### Input:\\n1, 1, 2, 3, 5, 8\\n\\n### Response:\\n13, 21, 34, 55, 89, 144, 233, 377, 610, 987, 1597, 2584, 4181, 6765, 10946, 17711, 28657, 46368, 75025']" + ] + }, + "execution_count": 9, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "# alpaca_prompt = Copied from above\n", + "FastLanguageModel.for_inference(model) # Enable native 2x faster inference\n", + "inputs = tokenizer(\n", + "[\n", + " alpaca_prompt.format(\n", + " \"Continue the fibonnaci sequence.\", # instruction\n", + " \"1, 1, 2, 3, 5, 8\", # input\n", + " \"\", # output - leave this blank for generation!\n", + " )\n", + "], return_tensors = \"pt\").to(\"cuda\")\n", + "\n", + "outputs = model.generate(**inputs, max_new_tokens = 64, use_cache = True)\n", + "tokenizer.batch_decode(outputs)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "CrSvZObor0lY" + }, + "source": [ + " You can also use a `TextStreamer` for continuous inference - so you can see the generation token by token, instead of waiting the whole time!" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "e2pEuRb1r2Vg", + "outputId": "b13f5e53-4ca4-4551-dffa-aaa3c514dca4" + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "<|begin_of_text|>Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.\n", + "\n", + "### Instruction:\n", + "Continue the fibonnaci sequence.\n", + "\n", + "### Input:\n", + "1, 1, 2, 3, 5, 8\n", + "\n", + "### Response:\n", + "13, 21, 34, 55, 89, 144<|end_of_text|>\n" + ] + } + ], + "source": [ + "# alpaca_prompt = Copied from above\n", + "FastLanguageModel.for_inference(model) # Enable native 2x faster inference\n", + "inputs = tokenizer(\n", + "[\n", + " alpaca_prompt.format(\n", + " \"Continue the fibonnaci sequence.\", # instruction\n", + " \"1, 1, 2, 3, 5, 8\", # input\n", + " \"\", # output - leave this blank for generation!\n", + " )\n", + "], return_tensors = \"pt\").to(\"cuda\")\n", + "\n", + "from transformers import TextStreamer\n", + "text_streamer = TextStreamer(tokenizer)\n", + "_ = model.generate(**inputs, streamer = text_streamer, max_new_tokens = 128)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "uMuVrWbjAzhc" + }, + "source": [ + "\n", + "### Saving, loading finetuned models\n", + "To save the final model as LoRA adapters, either use Huggingface's `push_to_hub` for an online save or `save_pretrained` for a local save.\n", + "\n", + "**[NOTE]** This ONLY saves the LoRA adapters, and not the full model. To save to 16bit or GGUF, scroll down!" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "upcOlWe7A1vc", + "outputId": "030a6e13-9371-4717-c5c5-d4e3563e0cca" + }, + "outputs": [ + { + "data": { + "text/plain": [ + "('lora_model/tokenizer_config.json',\n", + " 'lora_model/special_tokens_map.json',\n", + " 'lora_model/tokenizer.json')" + ] + }, + "execution_count": 11, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "model.save_pretrained(\"lora_model\") # Local saving\n", + "tokenizer.save_pretrained(\"lora_model\")\n", + "# model.push_to_hub(\"your_name/lora_model\", token = \"...\") # Online saving\n", + "# tokenizer.push_to_hub(\"your_name/lora_model\", token = \"...\") # Online saving" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "AEEcJ4qfC7Lp" + }, + "source": [ + "Now if you want to load the LoRA adapters we just saved for inference, set `False` to `True`:" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "colab": { + "base_uri": "https://localhost:8080/" + }, + "id": "MKX_XKs_BNZR", + "outputId": "f8e7d3fe-8e4d-49ee-944f-08e70cdc1d87" + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "<|begin_of_text|>Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request.\n", + "\n", + "### Instruction:\n", + "What is a famous tall tower in Paris?\n", + "\n", + "### Input:\n", + "\n", + "\n", + "### Response:\n", + "One of the most famous and iconic tall towers in Paris is the Eiffel Tower. Standing at 324 meters (1,063 feet) tall, this wrought iron tower is a symbol of the city and a must-see attraction for tourists from all over the world.<|end_of_text|>\n" + ] + } + ], + "source": [ + "if False:\n", + " from unsloth import FastLanguageModel\n", + " model, tokenizer = FastLanguageModel.from_pretrained(\n", + " model_name = \"lora_model\", # YOUR MODEL YOU USED FOR TRAINING\n", + " max_seq_length = max_seq_length,\n", + " dtype = dtype,\n", + " load_in_4bit = load_in_4bit,\n", + " )\n", + " FastLanguageModel.for_inference(model) # Enable native 2x faster inference\n", + "\n", + "# alpaca_prompt = You MUST copy from above!\n", + "\n", + "inputs = tokenizer(\n", + "[\n", + " alpaca_prompt.format(\n", + " \"What is a famous tall tower in Paris?\", # instruction\n", + " \"\", # input\n", + " \"\", # output - leave this blank for generation!\n", + " )\n", + "], return_tensors = \"pt\").to(\"cuda\")\n", + "\n", + "from transformers import TextStreamer\n", + "text_streamer = TextStreamer(tokenizer)\n", + "_ = model.generate(**inputs, streamer = text_streamer, max_new_tokens = 128)" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "QQMjaNrjsU5_" + }, + "source": [ + "You can also use Hugging Face's `AutoModelForPeftCausalLM`. Only use this if you do not have `unsloth` installed. It can be hopelessly slow, since `4bit` model downloading is not supported, and Unsloth's **inference is 2x faster**." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "yFfaXG0WsQuE" + }, + "outputs": [], + "source": [ + "if False:\n", + " # I highly do NOT suggest - use Unsloth if possible\n", + " from peft import AutoPeftModelForCausalLM\n", + " from transformers import AutoTokenizer\n", + " model = AutoPeftModelForCausalLM.from_pretrained(\n", + " \"lora_model\", # YOUR MODEL YOU USED FOR TRAINING\n", + " load_in_4bit = load_in_4bit,\n", + " )\n", + " tokenizer = AutoTokenizer.from_pretrained(\"lora_model\")" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "f422JgM9sdVT" + }, + "source": [ + "### Saving to float16 for VLLM\n", + "\n", + "We also support saving to `float16` directly. Select `merged_16bit` for float16 or `merged_4bit` for int4. We also allow `lora` adapters as a fallback. Use `push_to_hub_merged` to upload to your Hugging Face account! You can go to https://huggingface.co/settings/tokens for your personal tokens." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "iHjt_SMYsd3P" + }, + "outputs": [], + "source": [ + "# Merge to 16bit\n", + "if False: model.save_pretrained_merged(\"model\", tokenizer, save_method = \"merged_16bit\",)\n", + "if False: model.push_to_hub_merged(\"hf/model\", tokenizer, save_method = \"merged_16bit\", token = \"\")\n", + "\n", + "# Merge to 4bit\n", + "if False: model.save_pretrained_merged(\"model\", tokenizer, save_method = \"merged_4bit\",)\n", + "if False: model.push_to_hub_merged(\"hf/model\", tokenizer, save_method = \"merged_4bit\", token = \"\")\n", + "\n", + "# Just LoRA adapters\n", + "if False:\n", + " model.save_pretrained(\"model\")\n", + " tokenizer.save_pretrained(\"model\")\n", + "if False:\n", + " model.push_to_hub(\"hf/model\", token = \"\")\n", + " tokenizer.push_to_hub(\"hf/model\", token = \"\")\n" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "TCv4vXHd61i7" + }, + "source": [ + "### GGUF / llama.cpp Conversion\n", + "To save to `GGUF` / `llama.cpp`, we support it natively now! We clone `llama.cpp` and we default save it to `q8_0`. We allow all methods like `q4_k_m`. Use `save_pretrained_gguf` for local saving and `push_to_hub_gguf` for uploading to HF.\n", + "\n", + "Some supported quant methods (full list on our [Wiki page](https://github.com/unslothai/unsloth/wiki#gguf-quantization-options)):\n", + "* `q8_0` - Fast conversion. High resource use, but generally acceptable.\n", + "* `q4_k_m` - Recommended. Uses Q6_K for half of the attention.wv and feed_forward.w2 tensors, else Q4_K.\n", + "* `q5_k_m` - Recommended. Uses Q6_K for half of the attention.wv and feed_forward.w2 tensors, else Q5_K.\n", + "\n", + "[**NEW**] To finetune and auto export to Ollama, try our [Ollama notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3_(8B)-Ollama.ipynb)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": { + "id": "FqfebeAdT073" + }, + "outputs": [], + "source": [ + "# Save to 8bit Q8_0\n", + "if False: model.save_pretrained_gguf(\"model\", tokenizer,)\n", + "# Remember to go to https://huggingface.co/settings/tokens for a token!\n", + "# And change hf to your username!\n", + "if False: model.push_to_hub_gguf(\"hf/model\", tokenizer, token = \"\")\n", + "\n", + "# Save to 16bit GGUF\n", + "if False: model.save_pretrained_gguf(\"model\", tokenizer, quantization_method = \"f16\")\n", + "if False: model.push_to_hub_gguf(\"hf/model\", tokenizer, quantization_method = \"f16\", token = \"\")\n", + "\n", + "# Save to q4_k_m GGUF\n", + "if False: model.save_pretrained_gguf(\"model\", tokenizer, quantization_method = \"q4_k_m\")\n", + "if False: model.push_to_hub_gguf(\"hf/model\", tokenizer, quantization_method = \"q4_k_m\", token = \"\")\n", + "\n", + "# Save to multiple GGUF options - much faster if you want multiple!\n", + "if False:\n", + " model.push_to_hub_gguf(\n", + " \"hf/model\", # Change hf to your username!\n", + " tokenizer,\n", + " quantization_method = [\"q4_k_m\", \"q8_0\", \"q5_k_m\",],\n", + " token = \"\",\n", + " )" + ] + }, + { + "cell_type": "markdown", + "metadata": { + "id": "kGbBSRn6QKX7" + }, + "source": [ + "Now, use the `model-unsloth.gguf` file or `model-unsloth-Q4_K_M.gguf` file in llama.cpp.\n", + "\n", + "And we're done! If you have any questions on Unsloth, we have a [Discord](https://discord.gg/unsloth) channel! If you find any bugs or want to keep updated with the latest LLM stuff, or need help, join projects etc, feel free to join our Discord!\n", + "\n", + "Some other links:\n", + "1. Train your own reasoning model - Llama GRPO notebook [Free Colab](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.1_(8B)-GRPO.ipynb)\n", + "2. Saving finetunes to Ollama. [Free notebook](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3_(8B)-Ollama.ipynb)\n", + "3. Llama 3.2 Vision finetuning - Radiography use case. [Free Colab](https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.2_(11B)-Vision.ipynb)\n", + "6. See notebooks for DPO, ORPO, Continued pretraining, conversational finetuning and more on our [documentation](https://docs.unsloth.ai/get-started/unsloth-notebooks)!\n", + "\n", + "

\n", + " \n", + " \n", + " \n", + "\n", + " Join Discord if you need help + ⭐️ Star us on Github ⭐️\n", + "\n", + " This notebook and all Unsloth notebooks are licensed [LGPL-3.0](https://github.com/unslothai/notebooks?tab=LGPL-3.0-1-ov-file#readme).\n", + "
\n" + ] + } + ], + "metadata": { + "accelerator": "GPU", + "colab": { + "gpuType": "T4", + "provenance": [] + }, + "kernelspec": { + "display_name": "Python 3", + "name": "python3" + }, + "language_info": { + "name": "python" + }, + "widgets": { + "application/vnd.jupyter.widget-state+json": { + "0078f897f2174217a307d95d4f9bd775": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_8c195b5809604905b5e404baa30e8449", + "placeholder": "​", + "style": "IPY_MODEL_2247efac4283489bbd228330344388ab", + "value": "Map: 100%" + } + }, + "00d425bca350451da6400f9f05c4a659": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "03f492b4b56f4d8e80e9395a65058b1b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_9695a640b0ff4e91af495bb59548e4b6", + "placeholder": "​", + "style": "IPY_MODEL_d4bd5559d4134d64a943d57972c6ef39", + "value": "Map (num_proc=2): 100%" + } + }, + "04cc963133d242779572d2e847fa3d65": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "07055fc12b0841aaa5317f8252b5d347": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "0a9dc233674e4096b7a988a5e4ebaf84": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_d4b3770433bc41818372b7aed243fb31", + "placeholder": "​", + "style": "IPY_MODEL_19bcefcc1d874840ae9a9ca983e474b6", + "value": "Downloading readme: 100%" + } + }, + "0c1835f404db4846bb13b5da8d8f4447": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "19bcefcc1d874840ae9a9ca983e474b6": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "1a6db9aea6a64ae3aaef51d6265b35b2": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "1bd75ddaf57c4438a4e2c3070b9cef65": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_a3a3ef6d6337403cabea8b23f7c3021b", + "IPY_MODEL_c2ea0a3f01f34ffa8c94ab9b5098e9da", + "IPY_MODEL_68ea1d7cb8274a639b3fb5326f4218c3" + ], + "layout": "IPY_MODEL_39fef7b257614a0595f39355fa226b69" + } + }, + "1c7bc5fdb7dd4c39af8d4c2c504ec3ed": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_843a27e619534ea8914f9d36386c364b", + "IPY_MODEL_8f5adc70fbf248f2811527f620553be5", + "IPY_MODEL_4ffb4b2f015046fb94c1115ed0397a20" + ], + "layout": "IPY_MODEL_5a232ed040f94633a2a374031284c1f6" + } + }, + "1eb90e686e214122ae763b1b79ae321d": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "2006be31c09349738e221295bb84939f": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "2247efac4283489bbd228330344388ab": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "29c5b713f07043dda51820523e5c8ff3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_d7375f0f048841b29a20601c122666e8", + "IPY_MODEL_f433ced9bfcd4a57ba691d3c1caeed08", + "IPY_MODEL_da8ffc70820a48f5a12c6d4b5967015b" + ], + "layout": "IPY_MODEL_1a6db9aea6a64ae3aaef51d6265b35b2" + } + }, + "2b848e5a85bc42bc87945fd9ed5db038": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_bb9f3379310d4b04be694996f3137b28", + "max": 230, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_75082ba15db445df907f5612976590ae", + "value": 230 + } + }, + "2be29a4553ad4dfea8a9bc620c81a3ae": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "331f516c7a76456d801bc2a2feb228aa": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "33843107b93647b28985bfc37ea781ca": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "3662f1445ef34a50b462e601ed31bb69": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "36799fbcd90d43128620ff98225a825d": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "3719bf6f9c6a4c6fbef93c5328c11a07": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_03f492b4b56f4d8e80e9395a65058b1b", + "IPY_MODEL_39d9ef9fb35f47119f319f48eb222070", + "IPY_MODEL_3d7cfb33ceaf417e851ac4393c65148b" + ], + "layout": "IPY_MODEL_ece66fa2f128456fa2a82b8a28d1211c" + } + }, + "374fa9beda4042e1bf9a9b13de6e6674": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_4011ce9370d74fad857ec8e1e99d314f", + "max": 11610, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_41f5fed060ad4c8d87b24602b720ef04", + "value": 11610 + } + }, + "39d9ef9fb35f47119f319f48eb222070": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_fbae6e599d1644f39e5d86efa0f9f997", + "max": 51760, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_00d425bca350451da6400f9f05c4a659", + "value": 51760 + } + }, + "39fef7b257614a0595f39355fa226b69": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "3d7cfb33ceaf417e851ac4393c65148b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_6a27d9ad4f064586a87636b10455d15b", + "placeholder": "​", + "style": "IPY_MODEL_77f4367616964a01a8c42416f5f4c147", + "value": " 51760/51760 [00:50<00:00, 1965.57 examples/s]" + } + }, + "3fef797403d14440afe599a3bf06b626": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "4011ce9370d74fad857ec8e1e99d314f": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "41f5fed060ad4c8d87b24602b720ef04": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "4228734651ca45e19fc7bda79817f9b3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_f878c2e00bc240c7b0333cce950080e1", + "placeholder": "​", + "style": "IPY_MODEL_6b908368de51428585552dfef6a83088", + "value": " 5.70G/5.70G [00:45<00:00, 645MB/s]" + } + }, + "4613edbbec6846edb5b1677c25d542b6": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "484e507f14424f2b9173595b985f4101": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "4e7cb8e988114ed4b6fe09ff9f682dff": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "4ffb4b2f015046fb94c1115ed0397a20": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_8653acb618ad4e76bbf1daa00ea71238", + "placeholder": "​", + "style": "IPY_MODEL_e3c3bd9c4c124b0a8c88c83c1fc747d3", + "value": " 50.6k/50.6k [00:00<00:00, 2.29MB/s]" + } + }, + "51ca174d26e94b5cb1e895aa3c770655": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "51cd9026b1664819a67712996ca97bd5": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "5310346dd579424fa676b8e8e64790e7": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "5a232ed040f94633a2a374031284c1f6": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "5e8825fb770b41529f2129113cebc4a9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_0a9dc233674e4096b7a988a5e4ebaf84", + "IPY_MODEL_374fa9beda4042e1bf9a9b13de6e6674", + "IPY_MODEL_e6533d3c91fd4359bc84ffd8e59af5a3" + ], + "layout": "IPY_MODEL_f9b01aebcbdc48a585b7942b0ee60a2d" + } + }, + "5f174718e5974a7cab024d113f662513": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "620c0de28ec74f71a021a2be96dccf3a": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "634ae4c6cfe04673b1cdc9c9cac4cbf9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "68a32b398e1c490393e01befdc260785": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "68c686291b50430faeef0de7840e2c4b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_f4519637bb43400a80ce83505101e8a5", + "placeholder": "​", + "style": "IPY_MODEL_805676b197c94f5aa45956daa354640b", + "value": "Downloading data: 100%" + } + }, + "68ea1d7cb8274a639b3fb5326f4218c3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_5310346dd579424fa676b8e8e64790e7", + "placeholder": "​", + "style": "IPY_MODEL_0c1835f404db4846bb13b5da8d8f4447", + "value": " 9.09M/9.09M [00:00<00:00, 17.1MB/s]" + } + }, + "6a27d9ad4f064586a87636b10455d15b": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "6b908368de51428585552dfef6a83088": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "6e1aff64771c402ab070f650562fa4c9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_0078f897f2174217a307d95d4f9bd775", + "IPY_MODEL_ecdeaab4f8c94d6dade63bb06857c969", + "IPY_MODEL_735b85f0a0e9411cac4d704a504fcfc1" + ], + "layout": "IPY_MODEL_8e992e60416145a8b6eed744287ca0fb" + } + }, + "6e3c281f112b4a86af7a3ef95933d221": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_92395f250a154006923aaf9ea0a9c30b", + "IPY_MODEL_f84bfc5390054ec687c157c4d68199a6", + "IPY_MODEL_4228734651ca45e19fc7bda79817f9b3" + ], + "layout": "IPY_MODEL_4613edbbec6846edb5b1677c25d542b6" + } + }, + "735b85f0a0e9411cac4d704a504fcfc1": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_e99423a1ed3f4f72886b39368468b7c1", + "placeholder": "​", + "style": "IPY_MODEL_5f174718e5974a7cab024d113f662513", + "value": " 51760/51760 [00:00<00:00, 52999.05 examples/s]" + } + }, + "75082ba15db445df907f5612976590ae": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "76e5410e286a4a5abd6c213a38aa38bb": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "77f4367616964a01a8c42416f5f4c147": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "7981edf408d54d41bbeac42da7492c6b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_76e5410e286a4a5abd6c213a38aa38bb", + "placeholder": "​", + "style": "IPY_MODEL_ae7e90a811f94e75997d6a9ed1be8596", + "value": "generation_config.json: 100%" + } + }, + "7a935956348e47c68fbdf05ddf4752f3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "805676b197c94f5aa45956daa354640b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "80a72037771e4da9be989eefabbc8e76": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_c97c40c2bf2a41a8ae1c75e0a9c8ebff", + "placeholder": "​", + "style": "IPY_MODEL_484e507f14424f2b9173595b985f4101", + "value": "Generating train split: 100%" + } + }, + "843a27e619534ea8914f9d36386c364b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_2006be31c09349738e221295bb84939f", + "placeholder": "​", + "style": "IPY_MODEL_07055fc12b0841aaa5317f8252b5d347", + "value": "tokenizer_config.json: 100%" + } + }, + "85cc6f24cba54563acb5598f54fed7b9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_80a72037771e4da9be989eefabbc8e76", + "IPY_MODEL_ba68b274c50b44ec9e02642378d271a6", + "IPY_MODEL_f84d2fe4f1c24a34948755abf1f32b7f" + ], + "layout": "IPY_MODEL_86511967834f4484a5ec4af387b7d7a9" + } + }, + "86511967834f4484a5ec4af387b7d7a9": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "8653acb618ad4e76bbf1daa00ea71238": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "88fcd51819b5483c9ab22df7ef89ab64": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "89934c4f26834f15b9889ec36fee3b65": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "8c195b5809604905b5e404baa30e8449": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "8e992e60416145a8b6eed744287ca0fb": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "8ea52b105a7e44978caca33c0e7e815b": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "8f5adc70fbf248f2811527f620553be5": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_1eb90e686e214122ae763b1b79ae321d", + "max": 50570, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_7a935956348e47c68fbdf05ddf4752f3", + "value": 50570 + } + }, + "92395f250a154006923aaf9ea0a9c30b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_d032fe2ba5d647d99026fdade758c0cd", + "placeholder": "​", + "style": "IPY_MODEL_e9971d220fe24552a1e9aa299765cfb9", + "value": "model.safetensors: 100%" + } + }, + "94730f13e92a4c9aac35c2cfb21fc48c": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "94cbb87829d1486899e2ff6325c2ecdf": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "953625aa1e824f8a8d203197b316b302": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_ce9fcc5eff1f460d80b703a4ca32dad1", + "max": 44307561, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_3fef797403d14440afe599a3bf06b626", + "value": 44307561 + } + }, + "9695a640b0ff4e91af495bb59548e4b6": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "98a6716e7438429ea322adb3e3264f91": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_68c686291b50430faeef0de7840e2c4b", + "IPY_MODEL_953625aa1e824f8a8d203197b316b302", + "IPY_MODEL_f899a815142542219bde22ff792fb60c" + ], + "layout": "IPY_MODEL_51ca174d26e94b5cb1e895aa3c770655" + } + }, + "9be9074028da42d39d044a78393a861f": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "9d45b9a5de3e4cba9ac35ad2cb187f51": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "a14bc1c2130842568a5fde6698731e5f": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "a3a3ef6d6337403cabea8b23f7c3021b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_d125995cc0934239a01ba01b78529f21", + "placeholder": "​", + "style": "IPY_MODEL_634ae4c6cfe04673b1cdc9c9cac4cbf9", + "value": "tokenizer.json: 100%" + } + }, + "ac5eacaaee8346c080e54ea7a52648a4": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_7981edf408d54d41bbeac42da7492c6b", + "IPY_MODEL_2b848e5a85bc42bc87945fd9ed5db038", + "IPY_MODEL_e88c33f37d6849e0b1a6b41254104cb9" + ], + "layout": "IPY_MODEL_33843107b93647b28985bfc37ea781ca" + } + }, + "ae7e90a811f94e75997d6a9ed1be8596": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "ba68b274c50b44ec9e02642378d271a6": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_68a32b398e1c490393e01befdc260785", + "max": 51760, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_04cc963133d242779572d2e847fa3d65", + "value": 51760 + } + }, + "bb9f3379310d4b04be694996f3137b28": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "c2ea0a3f01f34ffa8c94ab9b5098e9da": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_d7f92e8332374313bee87ccd427446a4", + "max": 9085657, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_36799fbcd90d43128620ff98225a825d", + "value": 9085657 + } + }, + "c97c40c2bf2a41a8ae1c75e0a9c8ebff": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "ce9fcc5eff1f460d80b703a4ca32dad1": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "d032fe2ba5d647d99026fdade758c0cd": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "d125995cc0934239a01ba01b78529f21": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "d4b3770433bc41818372b7aed243fb31": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "d4bd5559d4134d64a943d57972c6ef39": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "d7375f0f048841b29a20601c122666e8": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_331f516c7a76456d801bc2a2feb228aa", + "placeholder": "​", + "style": "IPY_MODEL_9be9074028da42d39d044a78393a861f", + "value": "special_tokens_map.json: 100%" + } + }, + "d7f92e8332374313bee87ccd427446a4": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "da8ffc70820a48f5a12c6d4b5967015b": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_8ea52b105a7e44978caca33c0e7e815b", + "placeholder": "​", + "style": "IPY_MODEL_3662f1445ef34a50b462e601ed31bb69", + "value": " 345/345 [00:00<00:00, 23.9kB/s]" + } + }, + "e3c3bd9c4c124b0a8c88c83c1fc747d3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "e6533d3c91fd4359bc84ffd8e59af5a3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_88fcd51819b5483c9ab22df7ef89ab64", + "placeholder": "​", + "style": "IPY_MODEL_e6ffac074f1b476ba2ade11b37732af3", + "value": " 11.6k/11.6k [00:00<00:00, 81.5kB/s]" + } + }, + "e6ffac074f1b476ba2ade11b37732af3": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "e887160635cb4803b9f33845df615ec6": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "e88c33f37d6849e0b1a6b41254104cb9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_89934c4f26834f15b9889ec36fee3b65", + "placeholder": "​", + "style": "IPY_MODEL_e887160635cb4803b9f33845df615ec6", + "value": " 230/230 [00:00<00:00, 11.6kB/s]" + } + }, + "e93d063faf984cc4aa51462418d9b57e": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "e99423a1ed3f4f72886b39368468b7c1": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "e9971d220fe24552a1e9aa299765cfb9": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "DescriptionStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "DescriptionStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "description_width": "" + } + }, + "ea55293415ca48a4be97c2e1e4769122": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "ecdeaab4f8c94d6dade63bb06857c969": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_e93d063faf984cc4aa51462418d9b57e", + "max": 51760, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_9d45b9a5de3e4cba9ac35ad2cb187f51", + "value": 51760 + } + }, + "ece66fa2f128456fa2a82b8a28d1211c": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "f433ced9bfcd4a57ba691d3c1caeed08": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_51cd9026b1664819a67712996ca97bd5", + "max": 345, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_ea55293415ca48a4be97c2e1e4769122", + "value": 345 + } + }, + "f4519637bb43400a80ce83505101e8a5": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "f84bfc5390054ec687c157c4d68199a6": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "ProgressView", + "bar_style": "danger", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_2be29a4553ad4dfea8a9bc620c81a3ae", + "max": 5702746390, + "min": 0, + "orientation": "horizontal", + "style": "IPY_MODEL_94cbb87829d1486899e2ff6325c2ecdf", + "value": 5702745847 + } + }, + "f84d2fe4f1c24a34948755abf1f32b7f": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_94730f13e92a4c9aac35c2cfb21fc48c", + "placeholder": "​", + "style": "IPY_MODEL_620c0de28ec74f71a021a2be96dccf3a", + "value": " 51760/51760 [00:01<00:00, 52026.13 examples/s]" + } + }, + "f878c2e00bc240c7b0333cce950080e1": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "f899a815142542219bde22ff792fb60c": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "1.5.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "1.5.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "1.5.0", + "_view_name": "HTMLView", + "description": "", + "description_tooltip": null, + "layout": "IPY_MODEL_4e7cb8e988114ed4b6fe09ff9f682dff", + "placeholder": "​", + "style": "IPY_MODEL_a14bc1c2130842568a5fde6698731e5f", + "value": " 44.3M/44.3M [00:00<00:00, 87.2MB/s]" + } + }, + "f9b01aebcbdc48a585b7942b0ee60a2d": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "fbae6e599d1644f39e5d86efa0f9f997": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "1.2.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "1.2.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "1.2.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "overflow_x": null, + "overflow_y": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + } + } + } + }, + "nbformat": 4, + "nbformat_minor": 0 +} diff --git a/tests/test_asft.py b/tests/test_asft.py new file mode 100644 index 0000000000..7781565468 --- /dev/null +++ b/tests/test_asft.py @@ -0,0 +1,1499 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for ASFT (Anchored Supervised Fine-Tuning) loss module.""" + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +HAS_CUDA = torch.cuda.is_available() +if not HAS_CUDA: + pytest.skip("CUDA is required for ASFT tests", allow_module_level = True) +torch.set_default_device("cuda") + +from unsloth.losses.asft import ( + ASFTStreamingConfig, + effective_logits, + fast_cross_entropy_loss_per_token, + build_shift_labels, + get_reference_forward_callable, + compute_asft_loss, + _compute_kl_divergence, + _compute_dft_weights, + _compute_kl_seq_kv_cache, +) + + +# ----------------------------------------------------------------------------- +# Test Fixtures +# ----------------------------------------------------------------------------- + + +@pytest.fixture +def dummy_logits(): + """Create dummy logits tensor (B=2, T=4, V=8).""" + torch.manual_seed(42) + return torch.randn(2, 4, 8, requires_grad = True) + + +@pytest.fixture +def dummy_labels(): + """Create dummy labels tensor with some -100 values.""" + # Labels: [0, 1, 2, 3] and [4, 5, -100, -100] + return torch.tensor([[0, 1, 2, 3], [4, 5, -100, -100]], dtype = torch.long) + + +@pytest.fixture +def simple_model(): + """Create a simple model for testing.""" + + class SimpleModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + final_logit_softcapping = 0, + logit_scale = 0, + ) + self.embedding = nn.Embedding(16, 8) + self.linear = nn.Linear(8, 8) + + def forward(self, input_ids = None, **kwargs): + # Deterministic forward with gradients + embeddings = self.embedding(input_ids) + logits = self.linear(embeddings) + return SimpleNamespace(logits = logits) + + return SimpleModel() + + +# ----------------------------------------------------------------------------- +# A1) Test effective_logits +# ----------------------------------------------------------------------------- + + +class TestEffectiveLogits: + """Tests for effective_logits function.""" + + def test_no_transformation(self, dummy_logits): + """Test that no transformation is applied when softcapping/scaling are 0.""" + result = effective_logits(dummy_logits, logit_softcapping = 0, logit_scaling = 0) + # Should be close to original (converted to float32) + assert torch.allclose(result, dummy_logits.float(), atol = 1e-6) + + def test_logit_scaling(self, dummy_logits): + """Test logit scaling: t * x.""" + scale = 2.0 + result = effective_logits(dummy_logits, logit_scaling = scale) + expected = scale * dummy_logits.float() + assert torch.allclose(result, expected, atol = 1e-6) + + def test_logit_softcapping(self, dummy_logits): + """Test logit softcapping: t * tanh(x / t).""" + softcap = 30.0 + result = effective_logits(dummy_logits, logit_softcapping = softcap) + expected = softcap * torch.tanh(dummy_logits.float() / softcap) + assert torch.allclose(result, expected, atol = 1e-6) + + def test_both_transformations(self, dummy_logits): + """Test both scaling and softcapping together.""" + scale = 2.0 + softcap = 30.0 + result = effective_logits( + dummy_logits, logit_softcapping = softcap, logit_scaling = scale + ) + # Scaling first, then softcapping + x = scale * dummy_logits.float() + expected = softcap * torch.tanh(x / softcap) + assert torch.allclose(result, expected, atol = 1e-6) + + def test_reads_from_model_config(self): + """Test reading config from model.""" + model = SimpleNamespace( + config = SimpleNamespace( + final_logit_softcapping = 30.0, + logit_scale = 2.0, + ) + ) + logits = torch.randn(2, 4, 8) + result = effective_logits(logits, model) + # Should apply both transformations + x = 2.0 * logits.float() + expected = 30.0 * torch.tanh(x / 30.0) + assert torch.allclose(result, expected, atol = 1e-6) + + def test_reads_granite_logit_scaling(self): + """Test Granite logit scaling override.""" + model = SimpleNamespace( + config = SimpleNamespace( + model_type = "granite", + final_logit_softcapping = 0, + logit_scale = 2.0, + logit_scaling = 0, + logits_scaling = 16.0, + ) + ) + logits = torch.randn(2, 4, 8) + result = effective_logits(logits, model) + expected = (1.0 / 16.0) * logits.float() + assert torch.allclose(result, expected, atol = 1e-6) + + def test_reads_falcon_h1_logit_scaling(self): + """Test Falcon H1 logit scaling override.""" + model = SimpleNamespace( + config = SimpleNamespace( + model_type = "falcon_h1", + final_logit_softcapping = 0, + logit_scale = 2.0, + logit_scaling = 0, + lm_head_multiplier = 3.0, + ) + ) + logits = torch.randn(2, 4, 8) + result = effective_logits(logits, model) + expected = 3.0 * logits.float() + assert torch.allclose(result, expected, atol = 1e-6) + + +# ----------------------------------------------------------------------------- +# A2) Test fast_cross_entropy_loss_per_token +# ----------------------------------------------------------------------------- + + +class TestFastCrossEntropyLossPerToken: + """Tests for fast_cross_entropy_loss_per_token function.""" + + def test_basic_loss_computation(self, dummy_logits, dummy_labels): + """Test basic per-token CE loss computation.""" + losses, valid_mask = fast_cross_entropy_loss_per_token( + dummy_logits.detach(), dummy_labels + ) + + # Check shapes + batch, seq_len = dummy_labels.shape + assert losses.shape == (batch * seq_len,) + assert valid_mask.shape == (batch * seq_len,) + + # Check that valid_mask correctly identifies -100 positions + flat_labels = dummy_labels.view(-1) + expected_valid = flat_labels != -100 + assert torch.equal(valid_mask, expected_valid) + + def test_ignored_positions_have_zero_loss(self, dummy_logits, dummy_labels): + """Test that positions with label -100 have zero loss.""" + losses, valid_mask = fast_cross_entropy_loss_per_token( + dummy_logits.detach(), dummy_labels + ) + + # Loss at ignored positions should be 0 + assert torch.all(losses[~valid_mask] == 0) + + def test_valid_positions_have_nonzero_loss(self, dummy_logits, dummy_labels): + """Test that valid positions have non-zero loss.""" + losses, valid_mask = fast_cross_entropy_loss_per_token( + dummy_logits.detach(), dummy_labels + ) + + # At least some valid positions should have non-zero loss + assert torch.any(losses[valid_mask] > 0) + + def test_matches_pytorch_ce(self): + """Test that results match PyTorch CE loss.""" + torch.manual_seed(42) + logits = torch.randn(2, 4, 8) + labels = torch.tensor([[0, 1, 2, 3], [4, 5, 6, 7]], dtype = torch.long) + + losses, valid_mask = fast_cross_entropy_loss_per_token(logits, labels) + + # Compare with PyTorch + flat_logits = logits.view(-1, 8) + flat_labels = labels.view(-1) + pytorch_losses = F.cross_entropy(flat_logits, flat_labels, reduction = "none") + + # Should be close + assert torch.allclose(losses, pytorch_losses, atol = 1e-4) + + def test_respects_custom_ignore_index(self): + """Test that custom ignore_index is honored by the kernel wrapper.""" + torch.manual_seed(0) + logits = torch.randn(1, 4, 8) + labels = torch.tensor([[1, 2, 1, 3]], dtype = torch.long) + + losses, valid_mask = fast_cross_entropy_loss_per_token( + logits, labels, ignore_index = 1 + ) + + assert losses.shape == (4,) + assert torch.equal(valid_mask, torch.tensor([False, True, False, True])) + assert torch.all(losses[~valid_mask] == 0) + + +# ----------------------------------------------------------------------------- +# A3) Test build_shift_labels +# ----------------------------------------------------------------------------- + + +class TestBuildShiftLabels: + """Tests for build_shift_labels function.""" + + def test_basic_shift(self): + """Test basic label shifting.""" + labels = torch.tensor([[0, 1, 2, 3], [4, 5, 6, 7]], dtype = torch.long) + shift_labels = build_shift_labels(labels) + + # shift_labels[..., :-1] = labels[..., 1:] + # shift_labels[..., -1] = -100 + expected = torch.tensor([[1, 2, 3, -100], [5, 6, 7, -100]], dtype = torch.long) + assert torch.equal(shift_labels, expected) + + def test_preserves_ignore_index(self): + """Test that existing -100 values are preserved after shift.""" + labels = torch.tensor([[0, 1, -100, -100], [4, 5, 6, -100]], dtype = torch.long) + shift_labels = build_shift_labels(labels) + + # First row: [1, -100, -100, -100] + # Second row: [5, 6, -100, -100] + expected = torch.tensor( + [[1, -100, -100, -100], [5, 6, -100, -100]], dtype = torch.long + ) + assert torch.equal(shift_labels, expected) + + def test_with_packed_seq_lengths(self): + """Test shift labels with packed sequence boundary masking.""" + # Single row with packed sequences of lengths [2, 2] + labels = torch.tensor([[0, 1, 2, 3]], dtype = torch.long) + packed_seq_lengths = torch.tensor([2, 2], dtype = torch.int32) + + shift_labels = build_shift_labels(labels, packed_seq_lengths) + + # After shift: [1, 2, 3, -100] + # After boundary masking at positions 1 and 3: [1, -100, 3, -100] + # Actually boundary positions are cumsum - 1 = [1, 3] + # So positions 1 and 3 should be -100 + assert shift_labels[0, 1].item() == -100 # End of first sequence + assert shift_labels[0, 3].item() == -100 # End of second sequence (also last) + + +# ----------------------------------------------------------------------------- +# A4) Test get_reference_forward_callable +# ----------------------------------------------------------------------------- + + +class TestGetReferenceForwardCallable: + """Tests for get_reference_forward_callable function.""" + + def test_disable_adapter_policy(self, simple_model): + """Test disable_adapter policy when model has adapters.""" + # Mock disable_adapter + simple_model.disable_adapter = MagicMock() + simple_model.disable_adapter.__enter__ = MagicMock(return_value = None) + simple_model.disable_adapter.__exit__ = MagicMock(return_value = False) + + ref_forward = get_reference_forward_callable( + simple_model, reference_policy = "disable_adapter" + ) + + # Call the forward + input_ids = torch.tensor([[1, 2, 3, 4]]) + result = ref_forward(input_ids = input_ids) + + # Should have called disable_adapter + assert simple_model.disable_adapter.__enter__.called + + def test_frozen_copy_policy(self, simple_model): + """Test frozen_copy policy.""" + ref_forward = get_reference_forward_callable( + simple_model, reference_policy = "frozen_copy" + ) + + input_ids = torch.tensor([[1, 2, 3, 4]]) + result = ref_forward(input_ids = input_ids) + + # Should return logits + assert result.shape[0] == 1 # batch size + assert result.shape[1] == 4 # seq len + + def test_fallback_to_frozen_copy_without_adapters(self, simple_model): + """Test that disable_adapter falls back to frozen_copy when no adapters.""" + # Model without disable_adapter method + ref_forward = get_reference_forward_callable( + simple_model, reference_policy = "disable_adapter" + ) + + input_ids = torch.tensor([[1, 2, 3, 4]]) + result = ref_forward(input_ids = input_ids) + + # Should still work (uses frozen copy fallback) + assert result is not None + + def test_return_outputs_true(self, simple_model): + """Test returning full outputs when requested.""" + ref_forward = get_reference_forward_callable( + simple_model, reference_policy = "frozen_copy", return_outputs = True + ) + + input_ids = torch.tensor([[1, 2, 3, 4]]) + result = ref_forward(input_ids = input_ids) + + assert hasattr(result, "logits") + + +# ----------------------------------------------------------------------------- +# Test KL divergence computation +# ----------------------------------------------------------------------------- + + +class TestKLDivergence: + """Tests for KL divergence computation.""" + + def test_kl_direction(self): + """Test that KL is computed as KL(p_ref || p_cur).""" + torch.manual_seed(42) + cur_logits = torch.randn(4, 8) # (B*T, V) + ref_logits = torch.randn(4, 8) + + kl = _compute_kl_divergence(cur_logits, ref_logits, kl_direction = "forward") + + # KL should be non-negative + assert torch.all(kl >= -1e-6) # Allow small numerical errors + + def test_kl_zero_for_identical(self): + """Test that KL is zero when distributions are identical.""" + logits = torch.randn(4, 8) + + kl = _compute_kl_divergence(logits, logits.clone(), kl_direction = "forward") + + # Should be close to zero + assert torch.allclose(kl, torch.zeros_like(kl), atol = 1e-5) + + def test_kl_shape(self): + """Test KL output shape.""" + cur_logits = torch.randn(2, 4, 8) # (B, T, V) + ref_logits = torch.randn(2, 4, 8) + + kl = _compute_kl_divergence(cur_logits, ref_logits, kl_direction = "forward") + + # Should be flattened to (B*T,) + assert kl.shape == (8,) + + def test_kl_reverse_matches_manual(self): + """Test reverse KL matches manual computation.""" + torch.manual_seed(321) + cur_logits = torch.randn(2, 5) + ref_logits = torch.randn(2, 5) + + kl_reverse = _compute_kl_divergence( + cur_logits, ref_logits, kl_direction = "reverse" + ) + + cur_p = F.softmax(cur_logits, dim = -1) + ref_p = F.softmax(ref_logits, dim = -1) + manual = (cur_p * (cur_p.log() - ref_p.log())).sum(dim = -1) + + assert torch.allclose(kl_reverse, manual, atol = 1e-5) + + +# ----------------------------------------------------------------------------- +# Test DFT weights computation +# ----------------------------------------------------------------------------- + + +class TestDFTWeights: + """Tests for DFT weights computation.""" + + def test_dft_weights_are_probabilities(self, dummy_logits, dummy_labels): + """Test that DFT weights are valid probabilities.""" + flat_logits = dummy_logits.detach().view(-1, 8) + flat_labels = dummy_labels.view(-1) + + weights = _compute_dft_weights(flat_logits, flat_labels) + + # Weights should be in [0, 1] + assert torch.all(weights >= 0) + assert torch.all(weights <= 1) + + def test_dft_weights_are_detached(self, dummy_logits, dummy_labels): + """Test that DFT weights are detached (no gradients).""" + weights = _compute_dft_weights( + dummy_logits.detach().view(-1, 8), + dummy_labels.view(-1), + ) + + assert not weights.requires_grad + + def test_dft_weights_match_exp_neg_ce(self): + """Test exp(-CE) matches softmax-gather for DFT weights.""" + torch.manual_seed(123) + logits = torch.randn(2, 3, 7) + labels = torch.tensor([[1, 2, 3], [4, 5, 6]], dtype = torch.long) + + ce_losses, valid_mask = fast_cross_entropy_loss_per_token(logits, labels) + + weights_from_ce = _compute_dft_weights( + logits, + labels, + ce_losses = ce_losses, + valid_mask = valid_mask, + ) + weights_from_softmax = _compute_dft_weights(logits, labels) + + assert torch.allclose(weights_from_ce, weights_from_softmax, atol = 1e-4) + + +# ----------------------------------------------------------------------------- +# A5) Test compute_asft_loss +# ----------------------------------------------------------------------------- + + +class TestComputeASFTLoss: + """Tests for the main compute_asft_loss function.""" + + def test_sft_mode(self, simple_model): + """Test SFT mode computes standard CE.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + loss = compute_asft_loss(simple_model, inputs, asft_mode = "sft", kl_weight = 0.0) + + # Should return a scalar loss + assert loss.dim() == 0 + assert loss.requires_grad + + def test_sft_mode_granite_logit_scaling(self): + """Test Granite logit scaling in ASFT CE path.""" + + class GraniteModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + model_type = "granite", + final_logit_softcapping = 0, + logit_scale = 2.0, + logit_scaling = 0, + logits_scaling = 8.0, + ) + self.embedding = nn.Embedding(16, 8) + self.linear = nn.Linear(8, 8) + + def forward(self, input_ids = None, **kwargs): + embeddings = self.embedding(input_ids) + logits = self.linear(embeddings) + return SimpleNamespace(logits = logits) + + model = GraniteModel() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + captured = {} + + def fake_ce( + logits, labels, logit_softcapping = 0, logit_scaling = 0, ignore_index = -100 + ): + captured["logit_scaling"] = logit_scaling + batch, seq_len, _ = logits.shape + losses = torch.zeros(batch * seq_len, device = logits.device) + valid_mask = labels.view(-1) != ignore_index + return losses, valid_mask + + with patch( + "unsloth.losses.asft.fast_cross_entropy_loss_per_token", + side_effect = fake_ce, + ): + loss = compute_asft_loss(model, inputs, asft_mode = "sft", kl_weight = 0.0) + + assert captured["logit_scaling"] == pytest.approx(1.0 / 8.0) + assert loss.dim() == 0 + + def test_sft_mode_falcon_h1_logit_scaling(self): + """Test Falcon H1 logit scaling in ASFT CE path.""" + + class FalconH1Model(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + model_type = "falcon_h1", + final_logit_softcapping = 0, + logit_scale = 0, + logit_scaling = 0, + lm_head_multiplier = 3.0, + ) + self.embedding = nn.Embedding(16, 8) + self.linear = nn.Linear(8, 8) + + def forward(self, input_ids = None, **kwargs): + embeddings = self.embedding(input_ids) + logits = self.linear(embeddings) + return SimpleNamespace(logits = logits) + + model = FalconH1Model() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + captured = {} + + def fake_ce( + logits, labels, logit_softcapping = 0, logit_scaling = 0, ignore_index = -100 + ): + captured["logit_scaling"] = logit_scaling + batch, seq_len, _ = logits.shape + losses = torch.zeros(batch * seq_len, device = logits.device) + valid_mask = labels.view(-1) != ignore_index + return losses, valid_mask + + with patch( + "unsloth.losses.asft.fast_cross_entropy_loss_per_token", + side_effect = fake_ce, + ): + loss = compute_asft_loss(model, inputs, asft_mode = "sft", kl_weight = 0.0) + + assert captured["logit_scaling"] == pytest.approx(3.0) + assert loss.dim() == 0 + + def test_dft_mode(self, simple_model): + """Test DFT mode.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + loss = compute_asft_loss(simple_model, inputs, asft_mode = "dft", kl_weight = 0.0) + + assert loss.dim() == 0 + assert loss.requires_grad + + def test_dft_normalize_by_weights(self, simple_model): + """Test DFT normalization by weight sum.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + logits = simple_model(input_ids = inputs["input_ids"]).logits + shift_labels = build_shift_labels(inputs["labels"]) + valid_mask = shift_labels != -100 + ce_losses, _ = fast_cross_entropy_loss_per_token(logits, shift_labels) + ce_losses = ce_losses.view(shift_labels.shape) + dft_weights = _compute_dft_weights( + logits, + shift_labels, + ce_losses = ce_losses, + valid_mask = valid_mask, + ).view(shift_labels.shape) + token_loss = ce_losses * dft_weights + expected = token_loss[valid_mask].sum() / dft_weights[ + valid_mask + ].sum().clamp_min(1e-8) + + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "dft", + kl_weight = 0.0, + normalize_by = "weights", + ) + + assert torch.allclose(loss, expected, atol = 1e-5) + + def test_sft_kl_mode(self, simple_model): + """Test SFT+KL mode.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + ) + + assert loss.dim() == 0 + assert loss.requires_grad + + def test_asft_mode(self, simple_model): + """Test full ASFT mode.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "asft", + kl_weight = 0.1, + reference_policy = "frozen_copy", + ) + + assert loss.dim() == 0 + assert loss.requires_grad + + def test_return_outputs(self, simple_model): + """Test return_outputs=True.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + loss, outputs = compute_asft_loss( + simple_model, inputs, asft_mode = "sft", return_outputs = True + ) + + assert loss.dim() == 0 + assert hasattr(outputs, "logits") + + def test_handles_all_ignored_labels(self, simple_model): + """Test that all -100 labels returns zero loss.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[-100, -100, -100, -100]]), + } + + loss = compute_asft_loss(simple_model, inputs, asft_mode = "sft") + + # Should return zero loss + assert loss.item() == 0.0 + + def test_uses_num_items_in_batch(self, simple_model): + """Test that num_items_in_batch is used for normalization.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + "num_items_in_batch": 2, # Override default + } + + loss = compute_asft_loss(simple_model, inputs, asft_mode = "sft") + + # Should use the provided n_items + assert loss.dim() == 0 + + def test_packing_boundary_masking(self, simple_model): + """Test that packed sequence boundaries are masked.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + "packed_seq_lengths": torch.tensor([2, 2], dtype = torch.int32), + } + + loss = compute_asft_loss(simple_model, inputs, asft_mode = "sft") + + # Should handle packing without error + assert loss.dim() == 0 + + +# ----------------------------------------------------------------------------- +# Test ASFTStreamingConfig +# ----------------------------------------------------------------------------- + + +class TestASFTStreamingConfig: + """Tests for ASFTStreamingConfig dataclass.""" + + def test_default_values(self): + """Test default configuration values.""" + config = ASFTStreamingConfig() + + assert config.mode is None + assert config.enabled is False + assert config.ref_strategy == "none" + assert config.ref_microbatch_size is None + assert config.seq_chunk_size is None + assert config.kl_token_chunk_size is None + assert config.force_fp32_kl is True + + def test_custom_values(self): + """Test custom configuration values.""" + config = ASFTStreamingConfig( + mode = "batch", + enabled = True, + ref_strategy = "batch_micro", + ref_microbatch_size = 4, + seq_chunk_size = 256, + ) + + assert config.mode == "batch" + assert config.enabled is True + assert config.ref_strategy == "batch_micro" + assert config.ref_microbatch_size == 4 + assert config.seq_chunk_size == 256 + + +class TestStreamingModeMapping: + """Tests for streaming mode routing in compute_asft_loss.""" + + def test_mode_batch_uses_batch_micro(self, simple_model): + """Test that mode=batch routes to batch micro.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4], [2, 3, 4, 5]]), + "labels": torch.tensor([[1, 2, 3, 4], [2, 3, 4, 5]]), + } + config = ASFTStreamingConfig( + mode = "batch", + ref_microbatch_size = 1, + enabled = False, + ref_strategy = "seq_kv_cache", + ) + + def batch_side_effect( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + microbatch_size, + logit_softcapping = 0, + logit_scaling = 0, + force_fp32 = True, + kl_direction = "forward", + ): + batch, seq_len = shift_labels.shape + return torch.zeros(batch, seq_len, device = shift_labels.device) + + with ( + patch( + "unsloth.losses.asft._compute_kl_batch_micro", + side_effect = batch_side_effect, + ) as batch_mock, + patch( + "unsloth.losses.asft._compute_kl_seq_kv_cache", + side_effect = AssertionError("seq_kv_cache should not be used"), + ), + ): + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = config, + ) + + assert batch_mock.called + assert batch_mock.call_args[0][6] == 1 + assert loss.dim() == 0 + + @pytest.mark.parametrize("mode", ["seq", "auto"]) + def test_mode_seq_and_auto_use_seq_kv_cache(self, mode, simple_model): + """Test that mode=seq/auto routes to seq_kv_cache.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + config = ASFTStreamingConfig( + mode = mode, + seq_chunk_size = 2, + enabled = False, + ref_strategy = "batch_micro", + ) + + def seq_side_effect( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size, + **kwargs, + ): + batch, seq_len = shift_labels.shape + return torch.zeros(batch, seq_len, device = shift_labels.device) + + with ( + patch( + "unsloth.losses.asft._compute_kl_seq_kv_cache", + side_effect = seq_side_effect, + ) as seq_mock, + patch( + "unsloth.losses.asft._compute_kl_batch_micro", + side_effect = AssertionError("batch_micro should not be used"), + ), + ): + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = config, + ) + + assert seq_mock.called + assert seq_mock.call_args[0][6] == 2 + assert seq_mock.call_args.kwargs["microbatch_size"] is None + assert loss.dim() == 0 + + def test_mode_hybrid_defaults_microbatch(self, simple_model): + """Test that hybrid mode sets a default microbatch size.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4], [2, 3, 4, 5]]), + "labels": torch.tensor([[1, 2, 3, 4], [2, 3, 4, 5]]), + } + config = ASFTStreamingConfig( + mode = "hybrid", + seq_chunk_size = 2, + ref_microbatch_size = None, + ) + + def seq_side_effect( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size, + **kwargs, + ): + batch, seq_len = shift_labels.shape + return torch.zeros(batch, seq_len, device = shift_labels.device) + + with patch( + "unsloth.losses.asft._compute_kl_seq_kv_cache", + side_effect = seq_side_effect, + ) as seq_mock: + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = config, + ) + + assert seq_mock.called + assert seq_mock.call_args.kwargs["microbatch_size"] == 1 + assert config.ref_microbatch_size is None + assert loss.dim() == 0 + + def test_mode_off_uses_full_forward(self, simple_model): + """Test that mode=off bypasses streaming helpers.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + config = ASFTStreamingConfig( + mode = "off", + enabled = True, + ref_strategy = "seq_kv_cache", + ) + + def kl_side_effect( + cur_logits, + ref_logits, + model = None, + logit_softcapping = 0, + logit_scaling = 0, + force_fp32 = True, + kl_direction = "forward", + ): + batch, seq_len = ref_logits.shape[:2] + return torch.zeros(batch * seq_len, device = ref_logits.device) + + with ( + patch( + "unsloth.losses.asft._compute_kl_divergence", + side_effect = kl_side_effect, + ) as kl_mock, + patch( + "unsloth.losses.asft._compute_kl_seq_kv_cache", + side_effect = AssertionError("seq_kv_cache should not be used"), + ), + patch( + "unsloth.losses.asft._compute_kl_batch_micro", + side_effect = AssertionError("batch_micro should not be used"), + ), + ): + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = config, + ) + + assert kl_mock.called + assert loss.dim() == 0 + + def test_invalid_mode_raises(self, simple_model): + """Test that invalid streaming mode raises a ValueError.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + config = ASFTStreamingConfig(mode = "invalid") + + with pytest.raises(ValueError): + compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = config, + ) + + +class TestSeqKVCacheStreaming: + """Tests for seq_kv_cache streaming behavior.""" + + def test_seq_kv_cache_runs_when_use_cache_false(self): + """Test that seq_kv_cache attempts chunking even if config.use_cache=False.""" + batch_size, seq_len, vocab_size = 1, 6, 5 + cur_logits = torch.randn(batch_size, seq_len, vocab_size) + shift_labels = torch.zeros(batch_size, seq_len, dtype = torch.long) + valid_mask = shift_labels != -100 + input_ids = torch.arange(seq_len).view(1, -1) + + class DummyModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = False, + final_logit_softcapping = 0, + logit_scale = 0, + ) + + model = DummyModel() + call_state = {"saw_past": False} + + def ref_forward(**kwargs): + input_ids_local = kwargs["input_ids"] + if input_ids_local.shape[1] == seq_len: + raise AssertionError("full forward not expected") + if "past_key_values" in kwargs: + call_state["saw_past"] = True + batch, chunk_len = input_ids_local.shape + logits = torch.zeros( + batch, chunk_len, vocab_size, device = input_ids_local.device + ) + return (logits, ("cache",)) + + forward_inputs = {"input_ids": input_ids} + + kl = _compute_kl_seq_kv_cache( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size = 4, + ) + + assert kl.shape == (batch_size, seq_len) + assert call_state["saw_past"] is True + + def test_seq_kv_cache_supports_microbatching(self): + """Test that seq_kv_cache can be microbatched by batch dimension.""" + batch_size, seq_len, vocab_size = 2, 4, 3 + cur_logits = torch.randn(batch_size, seq_len, vocab_size) + shift_labels = torch.zeros(batch_size, seq_len, dtype = torch.long) + valid_mask = shift_labels != -100 + input_ids = torch.arange(seq_len).repeat(batch_size, 1) + + class DummyModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = True, + final_logit_softcapping = 0, + logit_scale = 0, + ) + + model = DummyModel() + call_state = {"max_batch": 0} + + def ref_forward(**kwargs): + input_ids_local = kwargs["input_ids"] + call_state["max_batch"] = max( + call_state["max_batch"], input_ids_local.shape[0] + ) + if input_ids_local.shape[0] > 1: + raise AssertionError("expected microbatching") + batch, chunk_len = input_ids_local.shape + logits = torch.zeros( + batch, chunk_len, vocab_size, device = input_ids_local.device + ) + return (logits, ("cache",)) + + forward_inputs = {"input_ids": input_ids} + + kl = _compute_kl_seq_kv_cache( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size = 2, + microbatch_size = 1, + ) + + assert kl.shape == (batch_size, seq_len) + assert call_state["max_batch"] == 1 + + def test_seq_kv_cache_falls_back_to_batch_micro(self): + """Test that seq_kv_cache falls back to batch micro on cache failure.""" + batch_size, seq_len, vocab_size = 2, 6, 5 + cur_logits = torch.randn(batch_size, seq_len, vocab_size) + shift_labels = torch.zeros(batch_size, seq_len, dtype = torch.long) + valid_mask = shift_labels != -100 + input_ids = torch.arange(seq_len).repeat(batch_size, 1) + + class DummyModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = True, + final_logit_softcapping = 0, + logit_scale = 0, + ) + + model = DummyModel() + + def ref_forward(**kwargs): + input_ids_local = kwargs["input_ids"] + if ( + input_ids_local.shape[0] == batch_size + and input_ids_local.shape[1] == seq_len + ): + raise AssertionError("full forward not expected on fallback") + batch, chunk_len = input_ids_local.shape + logits = torch.zeros( + batch, chunk_len, vocab_size, device = input_ids_local.device + ) + return (logits, None) + + forward_inputs = {"input_ids": input_ids} + + kl = _compute_kl_seq_kv_cache( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size = 2, + ) + + assert kl.shape == (batch_size, seq_len) + + def test_seq_kv_cache_falls_back_with_packed_sequences(self): + """Test that packed sequences bypass seq_kv_cache chunking.""" + batch_size, seq_len, vocab_size = 2, 4, 3 + cur_logits = torch.randn(batch_size, seq_len, vocab_size) + shift_labels = torch.zeros(batch_size, seq_len, dtype = torch.long) + valid_mask = shift_labels != -100 + input_ids = torch.arange(seq_len).repeat(batch_size, 1) + packed_seq_lengths = torch.tensor([2, 2], dtype = torch.int32) + + class DummyModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = True, + final_logit_softcapping = 0, + logit_scale = 0, + ) + + model = DummyModel() + ref_forward = MagicMock() + forward_inputs = { + "input_ids": input_ids, + "packed_seq_lengths": packed_seq_lengths, + } + + def batch_side_effect( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + microbatch_size, + logit_softcapping = 0, + logit_scaling = 0, + force_fp32 = True, + kl_direction = "forward", + ): + batch, seq_len = shift_labels.shape + return torch.zeros(batch, seq_len, device = shift_labels.device) + + with patch( + "unsloth.losses.asft._compute_kl_batch_micro", + side_effect = batch_side_effect, + ) as batch_mock: + kl = _compute_kl_seq_kv_cache( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size = 2, + ) + + assert batch_mock.called + assert batch_mock.call_args[0][6] == 1 + assert not ref_forward.called + assert kl.shape == (batch_size, seq_len) + + def test_config_immutability_when_none_values(self, simple_model): + """Test that streaming_config is not mutated when values are None.""" + config = ASFTStreamingConfig( + enabled = True, + ref_strategy = "batch_micro", + ref_microbatch_size = None, # Should use default without mutation + ) + original_microbatch = config.ref_microbatch_size + original_chunk = config.seq_chunk_size + + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + # Call compute_asft_loss with sft mode (doesn't use streaming, but + # the config should still not be mutated) + loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft", + streaming_config = config, + ) + + # Config should not be mutated + assert config.ref_microbatch_size == original_microbatch + assert config.seq_chunk_size == original_chunk + + +# ----------------------------------------------------------------------------- +# Backward Compatibility Tests +# ----------------------------------------------------------------------------- + + +class TestBackwardCompatibility: + """Tests to ensure ASFT doesn't break existing behavior.""" + + def test_sft_mode_matches_standard_ce(self, simple_model): + """Test that SFT mode produces same results as standard CE.""" + torch.manual_seed(42) + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + # Compute ASFT loss in SFT mode + asft_loss = compute_asft_loss(simple_model, inputs, asft_mode = "sft") + + # The loss should be a valid scalar + assert asft_loss.dim() == 0 + assert not torch.isnan(asft_loss) + assert not torch.isinf(asft_loss) + + def test_streaming_equivalence(self, simple_model): + """Test that streaming produces equivalent results to full forward.""" + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + # Full forward + full_loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig(enabled = False), + ) + + # With batch_micro streaming (should be equivalent for batch=1) + streaming_loss = compute_asft_loss( + simple_model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig( + enabled = True, + ref_strategy = "batch_micro", + ref_microbatch_size = 1, + ), + ) + + # Should be very close + assert torch.allclose(full_loss, streaming_loss, atol = 1e-4) + + def test_seq_kv_cache_equivalence(self): + """Test that seq_kv_cache matches full forward for KL loss.""" + torch.manual_seed(123) + + class CacheModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = True, + final_logit_softcapping = 0, + logit_scale = 0, + ) + self.embedding = nn.Embedding(32, 8) + self.linear = nn.Linear(8, 32) + + def forward( + self, input_ids = None, past_key_values = None, use_cache = None, **kwargs + ): + embeddings = self.embedding(input_ids) + logits = self.linear(embeddings) + past = ( + ("cache",) if (use_cache or past_key_values is not None) else None + ) + return SimpleNamespace(logits = logits, past_key_values = past) + + model = CacheModel() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4, 5, 6], [6, 5, 4, 3, 2, 1]]), + "labels": torch.tensor([[1, 2, 3, 4, 5, 6], [6, 5, 4, 3, 2, 1]]), + } + + full_loss = compute_asft_loss( + model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig(enabled = False), + ) + + seq_loss = compute_asft_loss( + model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig( + enabled = True, + ref_strategy = "seq_kv_cache", + seq_chunk_size = 2, + ), + ) + + assert torch.allclose(full_loss, seq_loss, atol = 1e-4) + + def test_seq_kv_cache_microbatch_equivalence(self): + """Test that seq_kv_cache + microbatching matches full forward.""" + torch.manual_seed(456) + + class CacheModel(nn.Module): + def __init__(self): + super().__init__() + self.config = SimpleNamespace( + use_cache = True, + final_logit_softcapping = 0, + logit_scale = 0, + ) + self.embedding = nn.Embedding(32, 8) + self.linear = nn.Linear(8, 32) + + def forward( + self, input_ids = None, past_key_values = None, use_cache = None, **kwargs + ): + embeddings = self.embedding(input_ids) + logits = self.linear(embeddings) + past = ( + ("cache",) if (use_cache or past_key_values is not None) else None + ) + return SimpleNamespace(logits = logits, past_key_values = past) + + model = CacheModel() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4, 5, 6], [6, 5, 4, 3, 2, 1]]), + "labels": torch.tensor([[1, 2, 3, 4, 5, 6], [6, 5, 4, 3, 2, 1]]), + } + + full_loss = compute_asft_loss( + model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig(enabled = False), + ) + + combined_loss = compute_asft_loss( + model, + inputs, + asft_mode = "sft+kl", + kl_weight = 0.1, + reference_policy = "frozen_copy", + streaming_config = ASFTStreamingConfig( + enabled = True, + ref_strategy = "seq_kv_cache", + seq_chunk_size = 2, + ref_microbatch_size = 1, + ), + ) + + assert torch.allclose(full_loss, combined_loss, atol = 1e-4) + + +# ----------------------------------------------------------------------------- +# Integration Tests +# ----------------------------------------------------------------------------- + + +class TestASFTTrainerIntegration: + """Integration tests for ASFTTrainer.""" + + def test_import_asft_trainer(self): + """Test that ASFTTrainer can be imported.""" + from unsloth.trainer import ASFTTrainer, ASFTStreamingConfig + + assert ASFTTrainer is not None + assert ASFTStreamingConfig is not None + + def test_asft_trainer_inherits_unsloth_trainer(self): + """Test that ASFTTrainer inherits from UnslothTrainer.""" + from unsloth.trainer import ASFTTrainer, UnslothTrainer + + assert issubclass(ASFTTrainer, UnslothTrainer) + + +class TestASFTTrainerComputeLoss: + """Tests for ASFTTrainer.compute_loss behavior.""" + + def test_compute_loss_calls_asft_loss(self): + """Test ASFTTrainer compute_loss calls compute_asft_loss.""" + from unsloth.trainer import ASFTTrainer, ASFTStreamingConfig + + trainer = ASFTTrainer.__new__(ASFTTrainer) + trainer.asft_enabled = True + trainer.asft_mode = "sft" + trainer.kl_weight = 0.0 + trainer.kl_direction = "forward" + trainer.reference_policy = "disable_adapter" + trainer.asft_streaming = ASFTStreamingConfig() + trainer.normalize_by = "tokens" + trainer._asft_original_model = None + + model = nn.Module() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + expected = torch.tensor(1.0, device = inputs["input_ids"].device) + + with patch( + "unsloth.trainer.compute_asft_loss", return_value = expected + ) as loss_mock: + result = ASFTTrainer.compute_loss( + trainer, model, inputs, return_outputs = False, num_items_in_batch = 7 + ) + + assert result is expected + assert inputs["num_items_in_batch"] == 7 + assert loss_mock.called + assert loss_mock.call_args.kwargs["model"] is model + assert loss_mock.call_args.kwargs["asft_mode"] == "sft" + assert loss_mock.call_args.kwargs["kl_weight"] == 0.0 + assert loss_mock.call_args.kwargs["kl_direction"] == "forward" + assert loss_mock.call_args.kwargs["reference_policy"] == "disable_adapter" + assert loss_mock.call_args.kwargs["streaming_config"] is trainer.asft_streaming + assert loss_mock.call_args.kwargs["normalize_by"] == "tokens" + + def test_compute_loss_creates_frozen_copy_once(self): + """Test frozen copy is created once when needed.""" + from unsloth.trainer import ASFTTrainer, ASFTStreamingConfig + + trainer = ASFTTrainer.__new__(ASFTTrainer) + trainer.asft_enabled = True + trainer.asft_mode = "asft" + trainer.kl_weight = 0.1 + trainer.kl_direction = "forward" + trainer.reference_policy = "frozen_copy" + trainer.asft_streaming = ASFTStreamingConfig() + trainer.normalize_by = "tokens" + trainer._asft_original_model = None + + model = nn.Module() + model_copy = MagicMock() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + with ( + pytest.warns(UserWarning), + patch("unsloth.trainer.deepcopy", return_value = model_copy) as deepcopy_mock, + patch( + "unsloth.trainer.compute_asft_loss", + return_value = torch.tensor(0.5, device = inputs["input_ids"].device), + ), + ): + ASFTTrainer.compute_loss(trainer, model, inputs) + ASFTTrainer.compute_loss(trainer, model, inputs) + + assert deepcopy_mock.call_count == 1 + assert trainer._asft_original_model is model_copy + assert model_copy.eval.called + assert model_copy.requires_grad_.called + + def test_compute_loss_skips_copy_with_disable_adapter(self): + """Test disable_adapter policy skips frozen copy.""" + from unsloth.trainer import ASFTTrainer, ASFTStreamingConfig + + trainer = ASFTTrainer.__new__(ASFTTrainer) + trainer.asft_enabled = True + trainer.asft_mode = "asft" + trainer.kl_weight = 0.1 + trainer.kl_direction = "forward" + trainer.reference_policy = "disable_adapter" + trainer.asft_streaming = ASFTStreamingConfig() + trainer.normalize_by = "tokens" + trainer._asft_original_model = None + + model = MagicMock() + model.disable_adapter = MagicMock() + inputs = { + "input_ids": torch.tensor([[1, 2, 3, 4]]), + "labels": torch.tensor([[1, 2, 3, 4]]), + } + + with ( + patch("unsloth.trainer.deepcopy") as deepcopy_mock, + patch( + "unsloth.trainer.compute_asft_loss", + return_value = torch.tensor(0.5, device = inputs["input_ids"].device), + ) as loss_mock, + ): + ASFTTrainer.compute_loss(trainer, model, inputs) + + assert not deepcopy_mock.called + assert loss_mock.call_args.kwargs["original_model"] is None + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_unsloth_cli.py b/tests/test_unsloth_cli.py new file mode 100644 index 0000000000..97d998816f --- /dev/null +++ b/tests/test_unsloth_cli.py @@ -0,0 +1,80 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""CLI argument parsing tests for unsloth-cli.py.""" + +from pathlib import Path +import importlib.util + +import pytest + + +def _load_cli_module(): + root = Path(__file__).resolve().parents[1] + cli_path = root / "unsloth-cli.py" + spec = importlib.util.spec_from_file_location("unsloth_cli", cli_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def test_cli_defaults_asft(): + cli = _load_cli_module() + parser = cli.build_parser() + args = parser.parse_args([]) + + assert args.asft is False + assert args.asft_mode == "asft" + assert args.kl_weight == 0.0 + assert args.reference_policy == "disable_adapter" + assert args.asft_streaming == "off" + assert args.ref_microbatch_size is None + assert args.seq_chunk_size is None + + +def test_cli_asft_streaming_flag_defaults_auto(): + cli = _load_cli_module() + parser = cli.build_parser() + args = parser.parse_args(["--asft_streaming"]) + + assert args.asft_streaming == "auto" + + +def test_cli_asft_streaming_value(): + cli = _load_cli_module() + parser = cli.build_parser() + args = parser.parse_args(["--asft_streaming", "batch"]) + + assert args.asft_streaming == "batch" + + +def test_cli_asft_options_parsed(): + cli = _load_cli_module() + parser = cli.build_parser() + args = parser.parse_args( + [ + "--asft", + "--asft_mode", + "sft+kl", + "--kl_weight", + "0.2", + "--reference_policy", + "frozen_copy", + ] + ) + + assert args.asft is True + assert args.asft_mode == "sft+kl" + assert args.kl_weight == pytest.approx(0.2) + assert args.reference_policy == "frozen_copy" diff --git a/unsloth-cli.py b/unsloth-cli.py index 612da11eb2..f451ee082a 100644 --- a/unsloth-cli.py +++ b/unsloth-cli.py @@ -40,6 +40,7 @@ def run(args): from trl import SFTTrainer, SFTConfig from unsloth import is_bfloat16_supported from unsloth.models.loader_utils import prepare_device_map + from unsloth.trainer import ASFTTrainer, ASFTStreamingConfig import logging from unsloth import RawTextDataLoader @@ -155,13 +156,33 @@ def run(args): packing = args.packing, ) - # Initialize trainer - trainer = SFTTrainer( - model = model, - processing_class = tokenizer, - train_dataset = dataset, - args = training_args, - ) + # Initialize trainer - use ASFTTrainer if ASFT is enabled + asft_enabled = getattr(args, "asft", False) + if asft_enabled: + # Build ASFT streaming config + asft_streaming = ASFTStreamingConfig( + mode = getattr(args, "asft_streaming", None), + ref_microbatch_size = getattr(args, "ref_microbatch_size", None), + seq_chunk_size = getattr(args, "seq_chunk_size", None), + ) + trainer = ASFTTrainer( + model = model, + processing_class = tokenizer, + train_dataset = dataset, + args = training_args, + asft_enabled = True, + asft_mode = getattr(args, "asft_mode", "asft"), + kl_weight = getattr(args, "kl_weight", 0.0), + reference_policy = getattr(args, "reference_policy", "disable_adapter"), + asft_streaming = asft_streaming, + ) + else: + trainer = SFTTrainer( + model = model, + processing_class = tokenizer, + train_dataset = dataset, + args = training_args, + ) trainer.train() @@ -206,7 +227,7 @@ def run(args): print("Warning: The model is not saved!") -if __name__ == "__main__": +def build_parser(): parser = argparse.ArgumentParser( description = "🦥 Fine-tune your llm faster using unsloth!" ) @@ -469,5 +490,73 @@ if __name__ == "__main__": "--stride", type = int, default = 512, help = "Overlap between chunks" ) + # ASFT Options + asft_group = parser.add_argument_group( + "🎯 ASFT Options", + "Anchored Supervised Fine-Tuning loss configuration (off by default)", + ) + asft_group.add_argument( + "--asft", + action = "store_true", + help = "Enable ASFT (Anchored Supervised Fine-Tuning) loss computation", + ) + asft_group.add_argument( + "--asft_mode", + type = str, + default = "asft", + choices = ["sft", "dft", "sft+kl", "asft"], + help = ( + "ASFT loss mode: 'sft' (standard CE), 'dft' (CE weighted by confidence), " + "'sft+kl' (CE + KL from reference), 'asft' (DFT + KL). Default: 'asft'" + ), + ) + asft_group.add_argument( + "--kl_weight", + type = float, + default = 0.0, + help = "Weight for KL divergence term in sft+kl and asft modes. Default: 0.0", + ) + asft_group.add_argument( + "--reference_policy", + type = str, + default = "disable_adapter", + choices = ["disable_adapter", "frozen_copy"], + help = ( + "How to compute reference distribution: 'disable_adapter' (use model with LoRA disabled), " + "'frozen_copy' (use frozen deepcopy). Default: 'disable_adapter'" + ), + ) + asft_group.add_argument( + "--asft_streaming", + nargs = "?", + const = "auto", + default = "off", + choices = ["off", "auto", "batch", "seq", "hybrid"], + help = ( + "Streaming mode for reference forward: 'off' (full forward), " + "'auto' (seq_kv_cache with batch-micro fallback), " + "'batch' (microbatch by batch), 'seq' (sequence chunking with KV cache), " + "'hybrid' (batch micro + seq_kv_cache). " + "Use flag without value for 'auto'." + ), + ) + asft_group.add_argument( + "--ref_microbatch_size", + type = int, + default = None, + help = "Microbatch size for batch_micro or seq_kv_cache strategy", + ) + asft_group.add_argument( + "--seq_chunk_size", + type = int, + default = None, + help = "Sequence chunk size for seq_kv_cache strategy", + ) + + return parser + + +if __name__ == "__main__": + parser = build_parser() args = parser.parse_args() run(args) diff --git a/unsloth/losses/__init__.py b/unsloth/losses/__init__.py new file mode 100644 index 0000000000..70e96edd36 --- /dev/null +++ b/unsloth/losses/__init__.py @@ -0,0 +1,33 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Loss functions for Unsloth training.""" + +from .asft import ( + ASFTStreamingConfig, + compute_asft_loss, + effective_logits, + fast_cross_entropy_loss_per_token, + build_shift_labels, + get_reference_forward_callable, +) + +__all__ = [ + "ASFTStreamingConfig", + "compute_asft_loss", + "effective_logits", + "fast_cross_entropy_loss_per_token", + "build_shift_labels", + "get_reference_forward_callable", +] diff --git a/unsloth/losses/asft.py b/unsloth/losses/asft.py new file mode 100644 index 0000000000..e34883992c --- /dev/null +++ b/unsloth/losses/asft.py @@ -0,0 +1,1073 @@ +# Copyright 2023-present Daniel Han-Chen & the Unsloth team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Anchored Supervised Fine-Tuning (ASFT) Loss Module for Unsloth. + +This module implements ASFT as an optional loss-path for fine-tuning. +All logic is gated behind `asft_enabled` and does not modify existing +fast_cross_entropy_loss behavior. + +Unsloth also includes additional performance and VRAM optimizations around +the same ASFT objective. In this repository we refer to that practical, +optimized implementation as **ASFT+**. + +ASFT Modes: +- SFT: Standard cross-entropy loss +- DFT: CE weighted by model's confidence (detached) +- SFT+KL: CE + KL divergence from reference model +- ASFT: DFT + KL divergence (full ASFT loss) +""" + +from __future__ import annotations + +from contextlib import contextmanager +from copy import deepcopy +from dataclasses import dataclass, field +from typing import Any, Callable, Dict, Literal, Optional, Tuple, Union + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from unsloth.kernels.cross_entropy_loss import Fast_CrossEntropyLoss +from unsloth.utils.packing import mask_packed_sequence_boundaries + +__all__ = [ + "ASFTStreamingConfig", + "effective_logits", + "fast_cross_entropy_loss_per_token", + "build_shift_labels", + "get_reference_forward_callable", + "compute_asft_loss", +] + + +# Default chunk sizes for streaming strategies +_DEFAULT_SEQ_CHUNK_SIZE = 256 +_DEFAULT_REF_MICROBATCH_DIVISOR = 2 # batch_size // this value + + +# ----------------------------------------------------------------------------- +# Configuration +# ----------------------------------------------------------------------------- + + +@dataclass +class ASFTStreamingConfig: + """Configuration for ASFT streaming strategies to reduce VRAM peak. + + Attributes: + mode: High-level streaming mode. + - "off": Disable streaming. + - "auto": Try seq_kv_cache with automatic batch-micro fallback. + - "batch": Use batch microbatching only. + - "seq": Use seq_kv_cache (with fallback). + - "hybrid": Combine batch micro + seq_kv_cache. + enabled: Whether streaming is enabled. + ref_strategy: Strategy for reference model forward pass. + - "none": Full reference forward (no streaming). + - "batch_micro": Microbatch reference forward by batch dimension. + - "seq_kv_cache": Sequence chunking via KV cache. + ref_microbatch_size: Microbatch size for batch_micro or seq_kv_cache. + seq_chunk_size: Chunk size for seq_kv_cache strategy (e.g., 128-512). + kl_token_chunk_size: Optional extra chunking of valid tokens for KL. + force_fp32_kl: Whether to force FP32 for KL computation. + """ + + mode: Optional[Literal["off", "auto", "batch", "seq", "hybrid"]] = None + enabled: bool = False + ref_strategy: Literal["none", "batch_micro", "seq_kv_cache"] = "none" + ref_microbatch_size: Optional[int] = None + seq_chunk_size: Optional[int] = None + kl_token_chunk_size: Optional[int] = None + force_fp32_kl: bool = True + + +# ----------------------------------------------------------------------------- +# A1) Helper: effective_logits - Apply logit scaling and softcapping +# ----------------------------------------------------------------------------- + + +def _resolve_logit_params( + model: Optional[nn.Module], + logit_softcapping: Optional[float], + logit_scaling: Optional[float], +) -> Tuple[float, float]: + if model is not None: + config = getattr(model, "config", None) + if config is not None: + if logit_softcapping is None: + logit_softcapping = getattr(config, "final_logit_softcapping", 0) + if logit_softcapping is None: + logit_softcapping = 0 + if logit_scaling is None: + logit_scaling = getattr(config, "logit_scale", 0) + if logit_scaling is None: + logit_scaling = 0 + if logit_scaling == 0: + logit_scaling = getattr(config, "logit_scaling", 0) + if logit_scaling is None: + logit_scaling = 0 + model_type = getattr(config, "model_type", None) + if model_type == "granite": + logits_scaling = getattr(config, "logits_scaling", 1) + if logits_scaling is None: + logits_scaling = 1 + logit_scaling = 1 / logits_scaling + elif model_type == "falcon_h1": + logit_scaling = getattr(config, "lm_head_multiplier", 0) + if logit_scaling is None: + logit_scaling = 0 + + if logit_softcapping is None: + logit_softcapping = 0 + if logit_scaling is None: + logit_scaling = 0 + + return logit_softcapping, logit_scaling + + +def effective_logits( + logits: torch.Tensor, + model: Optional[nn.Module] = None, + logit_softcapping: Optional[float] = None, + logit_scaling: Optional[float] = None, +) -> torch.Tensor: + """Apply logit scaling and softcapping consistent with Unsloth's Triton CE kernel. + + This ensures DFT weights and KL are computed on the same distribution as CE. + + Args: + logits: Input logits tensor. + model: Model to extract config from (optional if softcapping/scaling provided). + logit_softcapping: Softcapping value (e.g., for Gemma 2). If None, read from model. + logit_scaling: Logit scaling value (e.g., for Cohere). If None, read from model. + + Returns: + Transformed logits with scaling and softcapping applied. + """ + logit_softcapping, logit_scaling = _resolve_logit_params( + model, + logit_softcapping, + logit_scaling, + ) + + # Convert to float32 for stability + x = logits.float() + + # Apply scaling: t * x + if logit_scaling != 0: + x = logit_scaling * x + + # Apply softcapping: t * tanh(x / t) + if logit_softcapping != 0: + x = logit_softcapping * torch.tanh(x / logit_softcapping) + + return x + + +# ----------------------------------------------------------------------------- +# A2) Helper: fast_cross_entropy_loss_per_token +# ----------------------------------------------------------------------------- + + +def fast_cross_entropy_loss_per_token( + logits: torch.Tensor, + labels: torch.Tensor, + logit_softcapping: float = 0, + logit_scaling: float = 0, + ignore_index: int = -100, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Compute per-token cross-entropy loss using Unsloth's Triton kernel. + + This is a wrapper around Fast_CrossEntropyLoss that returns token-level + losses without reduction, suitable for ASFT weighting. + + Args: + logits: Logits tensor of shape (B, T, V) or (B*T, V). + labels: Labels tensor of shape (B, T) or (B*T,). + logit_softcapping: Softcapping value for the kernel. + logit_scaling: Scaling value for the kernel. + ignore_index: Index to ignore in loss computation. + + Returns: + Tuple of: + - losses: Per-token losses of shape (B*T,) where ignored = 0. + - valid_mask: Boolean mask of shape (B*T,) indicating valid tokens. + """ + # Flatten if needed + original_shape = None + if logits.dim() == 3: + batch, seq_len, vocab_size = logits.shape + original_shape = (batch, seq_len) + logits = logits.view(batch * seq_len, vocab_size) + labels = labels.view(batch * seq_len) + else: + vocab_size = logits.shape[-1] + + # Create valid mask before computing loss + valid_mask = labels != ignore_index + + labels_for_kernel = labels + if ignore_index != -100: + labels_for_kernel = labels.clone() + labels_for_kernel[labels_for_kernel == ignore_index] = -100 + + # Compute per-token CE using Unsloth's Triton kernel + # The kernel already handles ignore_index (-100) internally and returns 0 for those + losses = Fast_CrossEntropyLoss.apply( + logits, + labels_for_kernel, + logit_softcapping, + logit_scaling, + ) + + return losses, valid_mask + + +# ----------------------------------------------------------------------------- +# A3) Helper: build_shift_labels - Unsloth-style label shifting +# ----------------------------------------------------------------------------- + + +def build_shift_labels( + labels: torch.Tensor, + packed_seq_lengths: Optional[torch.Tensor] = None, + ignore_index: int = -100, +) -> torch.Tensor: + """Build shifted labels in Unsloth style. + + Unsloth CE path "shifts labels, not logits" - the last token becomes -100. + + shift_labels[..., :-1] = labels[..., 1:] + shift_labels[..., -1] = -100 + + Args: + labels: Original labels tensor of shape (B, T). + packed_seq_lengths: Optional packed sequence lengths for boundary masking. + ignore_index: Index to use for ignored positions. + + Returns: + Shifted labels tensor of same shape as input. + """ + shift_labels = torch.empty_like(labels) + shift_labels[..., :-1] = labels[..., 1:] + shift_labels[..., -1] = ignore_index + + # Apply packing boundary masking if needed + if packed_seq_lengths is not None: + mask_packed_sequence_boundaries( + shift_labels, + packed_seq_lengths, + ignore_index = ignore_index, + ) + + return shift_labels + + +# ----------------------------------------------------------------------------- +# A4) Helper: get_reference_forward_callable +# ----------------------------------------------------------------------------- + + +@contextmanager +def _inference_eval_context(model: nn.Module): + """Context manager for inference mode with eval() state preserved.""" + was_training = model.training + try: + model.eval() + with torch.inference_mode(): + yield + finally: + if was_training: + model.train() + + +def get_reference_forward_callable( + model: nn.Module, + reference_policy: Literal["disable_adapter", "frozen_copy"] = "disable_adapter", + original_model: Optional[nn.Module] = None, + return_outputs: bool = False, +) -> Callable[..., torch.Tensor]: + """Get a callable for reference model forward pass. + + Args: + model: The main model (may have LoRA adapters). + reference_policy: How to get reference distribution: + - "disable_adapter": Use model with adapters disabled (requires PEFT). + - "frozen_copy": Use a frozen deepcopy of the model. + original_model: Optional pre-created frozen model for "frozen_copy" policy. + + Returns: + Callable that takes forward inputs (without labels) and returns logits. + """ + # Check for PEFT/LoRA adapters + has_adapters = hasattr(model, "disable_adapter") + + if reference_policy == "disable_adapter" and has_adapters: + # Use adapter-disabled model + def ref_forward(**forward_inputs) -> torch.Tensor: + with _inference_eval_context(model): + disable_adapter = model.disable_adapter + if hasattr(disable_adapter, "__enter__") and hasattr( + disable_adapter, "__exit__" + ): + context_manager = disable_adapter + else: + context_manager = disable_adapter() + with context_manager: + outputs = model(**forward_inputs) + return outputs if return_outputs else outputs.logits + + return ref_forward + + elif reference_policy == "frozen_copy" or not has_adapters: + # Use frozen copy + if original_model is None: + # Create frozen copy + original_model = deepcopy(model) + original_model.eval() + original_model.requires_grad_(False) + + def ref_forward(**forward_inputs) -> torch.Tensor: + with _inference_eval_context(original_model): + outputs = original_model(**forward_inputs) + return outputs if return_outputs else outputs.logits + + return ref_forward + + else: + raise ValueError(f"Unknown reference_policy: {reference_policy}") + + +# ----------------------------------------------------------------------------- +# Internal: KL divergence computation +# ----------------------------------------------------------------------------- + + +def _compute_kl_divergence( + cur_logits: torch.Tensor, + ref_logits: torch.Tensor, + model: Optional[nn.Module] = None, + logit_softcapping: float = 0, + logit_scaling: float = 0, + force_fp32: bool = True, + kl_direction: Literal["forward", "reverse"] = "forward", +) -> torch.Tensor: + """Compute per-token KL divergence (forward KL by default). + + KL naming is frequently confused due to PyTorch's kl_div signature + (target is the weighting distribution). Definitions here follow standard + math/RLHF convention: + - Forward KL: KL(p_ref || p_cur), expectation over p_ref (mass-covering). + - Reverse KL: KL(p_cur || p_ref), expectation over p_cur (mode-seeking). + The original ASFT repo's paper text says "reverse KL" but its code uses + F.kl_div(log(cur), ref), which is forward KL; we match the code behavior. + Reverse KL is available via kl_direction="reverse". + + Args: + cur_logits: Current model logits (B*T, V) or (B, T, V). + ref_logits: Reference model logits (same shape as cur_logits). + model: Model for extracting config (optional). + logit_softcapping: Softcapping value. + logit_scaling: Scaling value. + force_fp32: Whether to compute in FP32 for stability. + kl_direction: "forward" for KL(p_ref || p_cur), "reverse" for KL(p_cur || p_ref). + + Returns: + Per-token KL divergence of shape (B*T,) or (B, T). + """ + # Flatten if 3D + original_shape = None + if cur_logits.dim() == 3: + batch, seq_len, vocab_size = cur_logits.shape + original_shape = (batch, seq_len) + cur_logits = cur_logits.view(batch * seq_len, vocab_size) + ref_logits = ref_logits.view(batch * seq_len, vocab_size) + + # Apply effective logits transformation + cur_eff = effective_logits(cur_logits, model, logit_softcapping, logit_scaling) + ref_eff = effective_logits(ref_logits, model, logit_softcapping, logit_scaling) + + if force_fp32: + cur_eff = cur_eff.float() + ref_eff = ref_eff.float() + + if kl_direction == "forward": + # Forward KL: KL(p_ref || p_cur) + cur_logp = F.log_softmax(cur_eff, dim = -1) + ref_p = F.softmax(ref_eff, dim = -1) + # Using F.kl_div: kl_div(input=log_cur, target=ref) computes KL(ref || cur) + kl = F.kl_div(cur_logp, ref_p, reduction = "none").sum(dim = -1) + elif kl_direction == "reverse": + # Reverse KL: KL(p_cur || p_ref) + ref_logp = F.log_softmax(ref_eff, dim = -1) + cur_p = F.softmax(cur_eff, dim = -1) + kl = F.kl_div(ref_logp, cur_p, reduction = "none").sum(dim = -1) + else: + raise ValueError(f"Unknown kl_direction: {kl_direction}") + + return kl + + +# ----------------------------------------------------------------------------- +# Internal: DFT weight computation +# ----------------------------------------------------------------------------- + + +def _compute_dft_weights( + logits: torch.Tensor, + labels: torch.Tensor, + model: Optional[nn.Module] = None, + logit_softcapping: float = 0, + logit_scaling: float = 0, + ignore_index: int = -100, + ce_losses: Optional[torch.Tensor] = None, + valid_mask: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Compute DFT weights: probability of target token under current model. + + w = p(label) where p = softmax(effective_logits), detached. + If ce_losses is provided, compute weights via exp(-ce_losses) to avoid + a full softmax over the vocab. + + Args: + logits: Model logits (B*T, V) or (B, T, V). + labels: Labels tensor (B*T,) or (B, T). + model: Model for extracting config. + logit_softcapping: Softcapping value. + logit_scaling: Scaling value. + ignore_index: Index to ignore. + ce_losses: Optional per-token CE losses aligned with labels. + valid_mask: Optional mask of valid tokens (same shape as labels). + + Returns: + DFT weights of same shape as labels, detached. + """ + if ce_losses is not None: + weights = torch.exp(-ce_losses.detach()) + if valid_mask is not None: + weights = weights * valid_mask + return weights + + # Flatten if 3D + if logits.dim() == 3: + batch, seq_len, vocab_size = logits.shape + logits = logits.view(batch * seq_len, vocab_size) + labels = labels.view(batch * seq_len) + else: + vocab_size = logits.shape[-1] + + # Apply effective logits transformation + logits_eff = effective_logits(logits, model, logit_softcapping, logit_scaling) + + # Compute softmax probabilities + p = F.softmax(logits_eff.float(), dim = -1) + + # Safe labels for gather (clamp -100 to 0) + safe_labels = labels.clamp(min = 0, max = vocab_size - 1) + + # Gather probabilities at target positions + weights = p.gather(dim = -1, index = safe_labels.unsqueeze(-1)).squeeze(-1) + + # Detach - weights should not receive gradients + return weights.detach() + + +# ----------------------------------------------------------------------------- +# Streaming helpers +# ----------------------------------------------------------------------------- + + +def _unwrap_reference_outputs( + ref_outputs: Any, +) -> Tuple[torch.Tensor, Optional[Any]]: + """Extract logits and past_key_values from reference outputs.""" + if hasattr(ref_outputs, "logits"): + return ref_outputs.logits, getattr(ref_outputs, "past_key_values", None) + if isinstance(ref_outputs, (tuple, list)) and len(ref_outputs) > 0: + past_key_values = ref_outputs[1] if len(ref_outputs) > 1 else None + return ref_outputs[0], past_key_values + return ref_outputs, None + + +def _slice_batch_inputs( + forward_inputs: Dict[str, Any], + batch_size: int, + b_start: int, + b_end: int, +) -> Dict[str, Any]: + """Slice batch-first tensors for microbatch processing.""" + mb_inputs = {} + for key, value in forward_inputs.items(): + if torch.is_tensor(value) and value.shape[0] == batch_size: + mb_inputs[key] = value[b_start:b_end] + else: + mb_inputs[key] = value + return mb_inputs + + +def _compute_kl_batch_micro( + model: nn.Module, + cur_logits: torch.Tensor, + shift_labels: torch.Tensor, + valid_mask: torch.Tensor, + ref_forward: Callable, + forward_inputs: Dict[str, Any], + microbatch_size: int, + logit_softcapping: float = 0, + logit_scaling: float = 0, + force_fp32: bool = True, + kl_direction: Literal["forward", "reverse"] = "forward", +) -> torch.Tensor: + """Compute KL using batch microbatching strategy. + + Processes reference forward in microbatches to reduce peak VRAM. + + Args: + model: Current model. + cur_logits: Current model logits (B, T, V). + shift_labels: Shifted labels (B, T). + valid_mask: Valid token mask (B, T). + ref_forward: Reference forward callable. + forward_inputs: Forward inputs (without labels). + microbatch_size: Size of each microbatch. + logit_softcapping: Softcapping value. + logit_scaling: Scaling value. + force_fp32: Whether to use FP32 for KL. + kl_direction: "forward" for KL(p_ref || p_cur), "reverse" for KL(p_cur || p_ref). + + Returns: + KL tensor of shape (B, T). + """ + batch_size = cur_logits.shape[0] + device = cur_logits.device + kl = torch.zeros_like(shift_labels, dtype = torch.float32) + + for b_start in range(0, batch_size, microbatch_size): + b_end = min(b_start + microbatch_size, batch_size) + + # Slice inputs for microbatch + mb_inputs = _slice_batch_inputs(forward_inputs, batch_size, b_start, b_end) + + # Get reference logits for microbatch + ref_outputs_mb = ref_forward(**mb_inputs) + ref_logits_mb, _ = _unwrap_reference_outputs(ref_outputs_mb) + cur_logits_mb = cur_logits[b_start:b_end] + + # Compute KL for this microbatch + kl_mb = _compute_kl_divergence( + cur_logits_mb, + ref_logits_mb, + model, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + + # Reshape if needed + if kl_mb.dim() == 1: + mb_batch = b_end - b_start + kl_mb = kl_mb.view(mb_batch, -1) + + kl[b_start:b_end] = kl_mb + + # Free memory + del ref_logits_mb + + return kl + + +def _compute_kl_seq_kv_cache( + model: nn.Module, + cur_logits: torch.Tensor, + shift_labels: torch.Tensor, + valid_mask: torch.Tensor, + ref_forward: Callable, + forward_inputs: Dict[str, Any], + seq_chunk_size: int, + microbatch_size: Optional[int] = None, + allow_auto_microbatch_fallback: bool = True, + logit_softcapping: float = 0, + logit_scaling: float = 0, + force_fp32: bool = True, + kl_direction: Literal["forward", "reverse"] = "forward", +) -> torch.Tensor: + """Compute KL using sequence chunking with KV cache strategy. + + Processes reference forward in sequence chunks to reduce peak VRAM. + Falls back to full forward if model doesn't support caching. + + Args: + model: Current model. + cur_logits: Current model logits (B, T, V). + shift_labels: Shifted labels (B, T). + valid_mask: Valid token mask (B, T). + ref_forward: Reference forward callable. + forward_inputs: Forward inputs (without labels). + seq_chunk_size: Size of each sequence chunk. + microbatch_size: Optional microbatch size for batch dimension. + allow_auto_microbatch_fallback: Allow automatic microbatch fallback on errors. + logit_softcapping: Softcapping value. + logit_scaling: Scaling value. + force_fp32: Whether to use FP32 for KL. + kl_direction: "forward" for KL(p_ref || p_cur), "reverse" for KL(p_cur || p_ref). + + Returns: + KL tensor of shape (B, T). + """ + batch_size, seq_len, vocab_size = cur_logits.shape + device = cur_logits.device + + packed_seq_lengths = forward_inputs.get("packed_seq_lengths", None) + if packed_seq_lengths is not None: + # Avoid seq_kv_cache with packed sequences; fall back to batch/full reference. + fallback_microbatch = None + if microbatch_size is not None and microbatch_size < batch_size: + fallback_microbatch = microbatch_size + elif allow_auto_microbatch_fallback: + fallback_microbatch = max(1, batch_size // _DEFAULT_REF_MICROBATCH_DIVISOR) + if fallback_microbatch >= batch_size: + fallback_microbatch = None + if fallback_microbatch is not None: + return _compute_kl_batch_micro( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + fallback_microbatch, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + ref_outputs = ref_forward(**forward_inputs) + ref_logits, _ = _unwrap_reference_outputs(ref_outputs) + kl_full = _compute_kl_divergence( + cur_logits, + ref_logits, + model, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + if kl_full.dim() == 1: + kl_full = kl_full.view(batch_size, seq_len) + return kl_full + + if microbatch_size is not None: + microbatch_size = max(1, microbatch_size) + if microbatch_size is not None and microbatch_size < batch_size: + kl = torch.zeros(batch_size, seq_len, dtype = torch.float32, device = device) + for b_start in range(0, batch_size, microbatch_size): + b_end = min(b_start + microbatch_size, batch_size) + mb_inputs = _slice_batch_inputs(forward_inputs, batch_size, b_start, b_end) + kl_mb = _compute_kl_seq_kv_cache( + model, + cur_logits[b_start:b_end], + shift_labels[b_start:b_end], + valid_mask[b_start:b_end], + ref_forward, + mb_inputs, + seq_chunk_size, + microbatch_size = None, + allow_auto_microbatch_fallback = False, + logit_softcapping = logit_softcapping, + logit_scaling = logit_scaling, + force_fp32 = force_fp32, + kl_direction = kl_direction, + ) + if kl_mb.dim() == 1: + mb_batch = b_end - b_start + kl_mb = kl_mb.view(mb_batch, -1) + kl[b_start:b_end] = kl_mb + return kl + + kl = torch.zeros(batch_size, seq_len, dtype = torch.float32, device = device) + + # Process in chunks with KV cache + past_key_values = None + + for s_start in range(0, seq_len, seq_chunk_size): + s_end = min(s_start + seq_chunk_size, seq_len) + + # Build chunk inputs + chunk_inputs = {} + for key, value in forward_inputs.items(): + if key == "input_ids": + chunk_inputs[key] = value[:, s_start:s_end] + elif key == "attention_mask": + # For chunked processing, need attention mask up to s_end + chunk_inputs[key] = value[:, :s_end] + elif key == "position_ids": + chunk_inputs[key] = value[:, s_start:s_end] + elif ( + torch.is_tensor(value) + and value.dim() >= 2 + and value.shape[1] == seq_len + ): + chunk_inputs[key] = value[:, s_start:s_end] + else: + chunk_inputs[key] = value + + # Add past_key_values if available + if past_key_values is not None: + chunk_inputs["past_key_values"] = past_key_values + + chunk_inputs["use_cache"] = True + + try: + # Get reference logits for chunk + # Note: ref_forward may not support all these kwargs + ref_outputs = ref_forward(**chunk_inputs) + ref_logits_chunk, ref_past_key_values = _unwrap_reference_outputs( + ref_outputs + ) + if ref_past_key_values is None and s_end < seq_len: + # Can't continue without cache; fall back to batch micro if allowed + fallback_microbatch = None + if allow_auto_microbatch_fallback: + fallback_microbatch = ( + microbatch_size + if microbatch_size is not None + else max(1, batch_size // _DEFAULT_REF_MICROBATCH_DIVISOR) + ) + if fallback_microbatch is not None and fallback_microbatch < batch_size: + return _compute_kl_batch_micro( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + fallback_microbatch, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + ref_outputs = ref_forward(**forward_inputs) + ref_logits, _ = _unwrap_reference_outputs(ref_outputs) + kl_full = _compute_kl_divergence( + cur_logits, + ref_logits, + model, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + if kl_full.dim() == 1: + kl_full = kl_full.view(batch_size, seq_len) + return kl_full + past_key_values = ref_past_key_values + + cur_logits_chunk = cur_logits[:, s_start:s_end] + + # Compute KL for this chunk + kl_chunk = _compute_kl_divergence( + cur_logits_chunk, + ref_logits_chunk, + model, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + + if kl_chunk.dim() == 1: + chunk_len = s_end - s_start + kl_chunk = kl_chunk.view(batch_size, chunk_len) + + kl[:, s_start:s_end] = kl_chunk + + del ref_logits_chunk + + except (RuntimeError, ValueError, KeyError, TypeError) as e: + # Fallback to batch micro or full forward on KV cache errors + # These exceptions typically indicate the model doesn't support + # the chunked KV cache approach (e.g., missing past_key_values support) + fallback_microbatch = None + if allow_auto_microbatch_fallback: + fallback_microbatch = ( + microbatch_size + if microbatch_size is not None + else max(1, batch_size // _DEFAULT_REF_MICROBATCH_DIVISOR) + ) + if fallback_microbatch is not None and fallback_microbatch < batch_size: + return _compute_kl_batch_micro( + model, + cur_logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + fallback_microbatch, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + ref_outputs = ref_forward(**forward_inputs) + ref_logits, _ = _unwrap_reference_outputs(ref_outputs) + kl_full = _compute_kl_divergence( + cur_logits, + ref_logits, + model, + logit_softcapping, + logit_scaling, + force_fp32, + kl_direction, + ) + if kl_full.dim() == 1: + kl_full = kl_full.view(batch_size, seq_len) + return kl_full + + return kl + + +# ----------------------------------------------------------------------------- +# A5) Core: compute_asft_loss +# ----------------------------------------------------------------------------- + + +def compute_asft_loss( + model: nn.Module, + inputs: Dict[str, Any], + *, + asft_mode: Literal["sft", "dft", "sft+kl", "asft"] = "asft", + kl_weight: float = 0.0, + kl_direction: Literal["forward", "reverse"] = "forward", + reference_policy: Literal["disable_adapter", "frozen_copy"] = "disable_adapter", + streaming_config: Optional[ASFTStreamingConfig] = None, + original_model: Optional[nn.Module] = None, + normalize_by: Literal["tokens", "weights"] = "tokens", + return_outputs: bool = False, +) -> Union[torch.Tensor, Tuple[torch.Tensor, Any]]: + """Compute ASFT loss. + + This is the main entry point for ASFT loss computation. + + Args: + model: The model to train. + inputs: Input dictionary containing input_ids, labels, etc. + asft_mode: Loss mode: + - "sft": Standard CE loss + - "dft": CE weighted by model confidence + - "sft+kl": CE + KL divergence from reference + - "asft": DFT + KL divergence (full ASFT) + kl_weight: Weight for KL term (only for sft+kl and asft modes). + kl_direction: "forward" for KL(p_ref || p_cur), "reverse" for KL(p_cur || p_ref). + reference_policy: How to get reference distribution. + streaming_config: Configuration for streaming strategies. + original_model: Optional pre-created frozen reference model. + normalize_by: "tokens" (default, matches reference) or "weights" for DFT/ASFT. + return_outputs: Whether to return model outputs alongside loss. + + Returns: + Loss tensor, or (loss, outputs) tuple if return_outputs=True. + """ + if streaming_config is None: + streaming_config = ASFTStreamingConfig() + + # Resolve streaming mode (new API) vs legacy enabled/ref_strategy + mode = streaming_config.mode + if mode is None: + streaming_enabled = streaming_config.enabled + ref_strategy = streaming_config.ref_strategy + else: + if mode == "off": + streaming_enabled = False + ref_strategy = "none" + elif mode == "batch": + streaming_enabled = True + ref_strategy = "batch_micro" + elif mode in ("seq", "auto", "hybrid"): + streaming_enabled = True + ref_strategy = "seq_kv_cache" + else: + raise ValueError(f"Unknown streaming mode: {mode}") + + # Get model config for softcapping/scaling + logit_softcapping, logit_scaling = _resolve_logit_params( + model, + None, + None, + ) + + # Build forward inputs (without labels/num_items to force logits materialization) + forward_inputs = { + k: v for k, v in inputs.items() if k not in {"labels", "num_items_in_batch"} + } + + # Main forward pass - ASFT always needs logits + outputs = model(**forward_inputs) + logits = outputs.logits # (B, T, V) + + # Get labels and build shift_labels Unsloth-style + labels = inputs["labels"] + packed_seq_lengths = inputs.get("packed_seq_lengths", None) + shift_labels = build_shift_labels(labels, packed_seq_lengths) + + # Valid mask and normalization + valid_mask = shift_labels != -100 + n_items_tokens = inputs.get("num_items_in_batch", None) + if n_items_tokens is None: + n_items_tokens = valid_mask.sum() + n_items_tokens = max(n_items_tokens, 1) # Avoid division by zero + + # Handle edge case: no valid tokens + if valid_mask.sum() == 0: + zero_loss = logits.sum() * 0.0 + if return_outputs: + return zero_loss, outputs + return zero_loss + + # Compute per-token CE loss + ce_losses, _ = fast_cross_entropy_loss_per_token( + logits, shift_labels, logit_softcapping, logit_scaling + ) + # Reshape to (B, T) + batch_size, seq_len = shift_labels.shape + ce_losses = ce_losses.view(batch_size, seq_len) + + # Initialize token losses + dft_weights = None + if asft_mode == "sft": + # Standard SFT: just CE + token_loss = ce_losses + + elif asft_mode == "dft": + # DFT: CE weighted by model confidence + dft_weights = _compute_dft_weights( + logits, + shift_labels, + model, + logit_softcapping, + logit_scaling, + ce_losses = ce_losses, + valid_mask = valid_mask, + ) + dft_weights = dft_weights.view(batch_size, seq_len) + token_loss = ce_losses * dft_weights + + elif asft_mode in ("sft+kl", "asft"): + # Need KL divergence + needs_outputs = streaming_enabled and ref_strategy == "seq_kv_cache" + ref_forward = get_reference_forward_callable( + model, + reference_policy, + original_model, + return_outputs = needs_outputs, + ) + + # Compute KL based on streaming strategy + # Use local variables to avoid mutating the input config + if streaming_enabled and ref_strategy == "batch_micro": + ref_microbatch_size = streaming_config.ref_microbatch_size + if ref_microbatch_size is None: + ref_microbatch_size = max( + 1, batch_size // _DEFAULT_REF_MICROBATCH_DIVISOR + ) + kl = _compute_kl_batch_micro( + model, + logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + ref_microbatch_size, + logit_softcapping, + logit_scaling, + streaming_config.force_fp32_kl, + kl_direction, + ) + elif streaming_enabled and ref_strategy == "seq_kv_cache": + seq_chunk_size = streaming_config.seq_chunk_size + if seq_chunk_size is None: + seq_chunk_size = _DEFAULT_SEQ_CHUNK_SIZE + ref_microbatch_size = streaming_config.ref_microbatch_size + if mode == "hybrid" and ref_microbatch_size is None: + ref_microbatch_size = max( + 1, batch_size // _DEFAULT_REF_MICROBATCH_DIVISOR + ) + allow_auto_microbatch_fallback = True + kl = _compute_kl_seq_kv_cache( + model, + logits, + shift_labels, + valid_mask, + ref_forward, + forward_inputs, + seq_chunk_size, + microbatch_size = ref_microbatch_size, + allow_auto_microbatch_fallback = allow_auto_microbatch_fallback, + logit_softcapping = logit_softcapping, + logit_scaling = logit_scaling, + force_fp32 = streaming_config.force_fp32_kl, + kl_direction = kl_direction, + ) + else: + # Full reference forward + ref_outputs = ref_forward(**forward_inputs) + ref_logits, _ = _unwrap_reference_outputs(ref_outputs) + kl = _compute_kl_divergence( + logits, + ref_logits, + model, + logit_softcapping, + logit_scaling, + streaming_config.force_fp32_kl, + kl_direction, + ) + kl = kl.view(batch_size, seq_len) + del ref_logits + + if asft_mode == "sft+kl": + # SFT + KL + token_loss = ce_losses + kl_weight * kl + else: + # Full ASFT: DFT + KL + dft_weights = _compute_dft_weights( + logits, + shift_labels, + model, + logit_softcapping, + logit_scaling, + ce_losses = ce_losses, + valid_mask = valid_mask, + ) + dft_weights = dft_weights.view(batch_size, seq_len) + dft_loss = ce_losses * dft_weights + token_loss = dft_loss + kl_weight * kl + + else: + raise ValueError(f"Unknown asft_mode: {asft_mode}") + + # Final reduction: sum over valid tokens, divide by chosen normalizer. + normalizer = n_items_tokens + if normalize_by == "weights" and dft_weights is not None: + weight_sum = dft_weights[valid_mask].sum() + normalizer = weight_sum.clamp_min(1e-8) + elif normalize_by != "tokens": + raise ValueError(f"Unknown normalize_by: {normalize_by}") + loss = token_loss[valid_mask].sum() / normalizer + + if return_outputs: + return loss, outputs + return loss diff --git a/unsloth/trainer.py b/unsloth/trainer.py index 65abe6801f..5b283fc37f 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -17,7 +17,8 @@ import os import psutil import warnings from dataclasses import dataclass, field -from typing import Optional +from typing import Literal, Optional, Union +from copy import deepcopy from functools import wraps import trl @@ -40,9 +41,17 @@ from unsloth_zoo.hf_utils import get_transformers_model_type from unsloth_zoo.utils import Version import dataclasses +# Import ASFT components +from unsloth.losses.asft import ( + ASFTStreamingConfig, + compute_asft_loss, +) + __all__ = [ "UnslothTrainingArguments", "UnslothTrainer", + "ASFTTrainer", + "ASFTStreamingConfig", "unsloth_train", "_patch_trl_trainer", "UnslothVisionDataCollator", @@ -198,6 +207,130 @@ class UnslothTrainer(SFTTrainer): return self.optimizer +class ASFTTrainer(UnslothTrainer): + """Trainer with ASFT (Anchored Supervised Fine-Tuning) loss support. + + ASFT provides alternative loss functions that weight tokens based on + model confidence and/or maintain similarity to a reference model. + + When asft_enabled=False (default), this trainer behaves identically + to UnslothTrainer/SFTTrainer with no changes to loss computation. + + Attributes: + asft_enabled: Whether to use ASFT loss computation. + asft_mode: Loss mode ("sft", "dft", "sft+kl", "asft"). + kl_weight: Weight for KL divergence term. + reference_policy: How to get reference distribution. + asft_streaming: Streaming configuration for VRAM reduction. + """ + + def __init__( + self, + *args, + asft_enabled: bool = False, + asft_mode: Literal["sft", "dft", "sft+kl", "asft"] = "asft", + kl_weight: float = 0.0, + kl_direction: Literal["forward", "reverse"] = "forward", + reference_policy: Literal["disable_adapter", "frozen_copy"] = "disable_adapter", + asft_streaming: Optional[ASFTStreamingConfig] = None, + normalize_by: Literal["tokens", "weights"] = "tokens", + **kwargs, + ): + """Initialize ASFTTrainer. + + Args: + *args: Positional arguments for parent trainer. + asft_enabled: Whether to enable ASFT loss. Default False preserves + standard SFT behavior completely unchanged. + asft_mode: Loss computation mode: + - "sft": Standard cross-entropy (for debugging/comparison) + - "dft": CE weighted by model's token probability + - "sft+kl": CE + KL divergence from reference + - "asft": Full ASFT (DFT + KL) + kl_weight: Weight for KL term (used in sft+kl and asft modes). + kl_direction: "forward" for KL(p_ref || p_cur), "reverse" for KL(p_cur || p_ref). + reference_policy: How to compute reference distribution: + - "disable_adapter": Use model with LoRA adapters disabled + - "frozen_copy": Use a frozen deepcopy of the model + asft_streaming: Optional streaming config for VRAM reduction. + normalize_by: "tokens" (default) or "weights" for DFT/ASFT normalization. + **kwargs: Keyword arguments for parent trainer. + """ + super().__init__(*args, **kwargs) + + self.asft_enabled = asft_enabled + self.asft_mode = asft_mode + self.kl_weight = kl_weight + self.kl_direction = kl_direction + self.reference_policy = reference_policy + self.asft_streaming = asft_streaming or ASFTStreamingConfig() + self.normalize_by = normalize_by + + # Will be lazily initialized if needed + self._asft_original_model = None + + def compute_loss(self, model, inputs, return_outputs = False, **kwargs): + """Compute loss with optional ASFT path. + + When asft_enabled=False, delegates entirely to parent compute_loss. + When asft_enabled=True, uses ASFT loss computation. + + Args: + model: The model to compute loss for. + inputs: Input dictionary. + return_outputs: Whether to return model outputs. + **kwargs: Additional arguments. + + Returns: + Loss tensor, or (loss, outputs) tuple if return_outputs=True. + """ + # If ASFT is disabled, use standard path unchanged + if not self.asft_enabled: + return super().compute_loss( + model, inputs, return_outputs = return_outputs, **kwargs + ) + + num_items_in_batch = kwargs.get("num_items_in_batch") + if num_items_in_batch is not None: + inputs["num_items_in_batch"] = num_items_in_batch + + if self.asft_mode in ("sft+kl", "asft"): + needs_frozen_copy = self.reference_policy == "frozen_copy" or ( + self.reference_policy == "disable_adapter" + and not hasattr(model, "disable_adapter") + ) + if needs_frozen_copy and self._asft_original_model is None: + if self.reference_policy == "frozen_copy": + warnings.warn( + "Unsloth: Creating a frozen copy of the model for ASFT. " + "This doubles VRAM usage. Use 'disable_adapter' if using LoRA.", + stacklevel = 2, + ) + elif self.reference_policy == "disable_adapter": + warnings.warn( + "Unsloth: 'disable_adapter' is unavailable; falling back to a " + "frozen copy for ASFT. This doubles VRAM usage.", + stacklevel = 2, + ) + self._asft_original_model = deepcopy(model) + self._asft_original_model.eval() + self._asft_original_model.requires_grad_(False) + + # ASFT-enabled path + return compute_asft_loss( + model = model, + inputs = inputs, + asft_mode = self.asft_mode, + kl_weight = self.kl_weight, + kl_direction = self.kl_direction, + reference_policy = self.reference_policy, + streaming_config = self.asft_streaming, + original_model = self._asft_original_model, + normalize_by = self.normalize_by, + return_outputs = return_outputs, + ) + + # From `trl>=0.13.0`, they changed how to pass several params to the trainer # We need to patch to make the transition smooth def _resolve_trainer_params(trainer_class, init_fn):