Compare commits

..
28 Commits
Author SHA1 Message Date
Aleksander GrygierandGitHub 6de1b63473 allozaur/feat/chat slash commands (#26716)
* base : slash-command/misc foundation - model icon and focus-selector constants

* feat : slash-command picker and command parsing helpers

* refactor : wire command and @-mention pickers into the chat form

* ui : improve model selector keyboard navigation and load/dismiss

* feat: Unify markdown/raw-text rendering under one setting with migration

* fix: Misc fixes - tool-call subtitle, assistant wrap, progress guards

* feat: Clamp and style numeric settings inputs from registry bounds
2026-08-07 20:20:01 +02:00
TitaniumtownandGitHub f8e30266d2 sycl: coalesce the ssm_conv window loads (#26612)
test-backend-ops perf -o SSM_CONV on an Arc Pro B70, interleaved A/B against
master, 6 reps, us/run:

  ne_a=[515,3328,1,1] ne_b=[4,3328,1,1]   n_t=512     97.68 -> 52.95   1.85x
  ne_a=[937,8192,1,1] ne_b=[4,8192,1,1]   n_t=934    516.16 -> 276.13  1.87x
  ne_a=[4,3328,1,1]   ne_b=[4,3328,1,1]   n_t=1        2.73 -> 2.71    flat

llama-bench on qwen35 27B Q4_K - Medium (48 of its 64 blocks run ssm_conv),
-ngl 99 -fa 1 -ctk f16 -ctv f16, interleaved passes of r=3:

  -b 2048 -ub 2048  pp2048  1045.1 / 1043.5 / 1043.7 -> 1069.5 / 1066.3 / 1065.9  +2.2%
  -b 2048 -ub 512   pp2048   771.8 /  772.7          ->  785.5 /  786.6           +1.8%
  -b 2048 -ub 512   tg128     23.81 /  23.88         ->   23.87 /  23.86          flat
2026-08-07 21:09:32 +03:00
robertomeroniandGitHub a194a75b7e metal : fix NORM/RMS_NORM for row lengths that leave a partial simdgroup (#26708)
ggml_metal_op_norm sized the threadgroup with
`nth = std::min(nth, args.ne00_t)`, which can leave nth not a multiple of
the simdgroup size. The kernels finish their row reduction with a
cross-simdgroup step where each lane of the last simdgroup reads one
per-simdgroup partial sum out of shmem_f32:

    if (tiisg == 0) { shmem_f32[sgitg] = sumf; }
    threadgroup_barrier(mem_flags::mem_threadgroup);
    sumf = shmem_f32[tiisg];
    sumf = simd_sum(sumf);

When the last simdgroup is partial it has fewer lanes than the
threadgroup has simdgroups, so the tail of the partial sums is never
read and the row sum is too small. For ne00_t = 33 nth becomes 33: two
simdgroups, but only one lane in the second, so one of the two partial
sums is dropped. The mean and variance are then wrong for the whole row.

Round ne00_t up to a whole number of simdgroups instead. Rounding up
rather than dropping the clamp keeps the threadgroup as small as
possible: deleting the line would raise nth to the next power of two
(ne00_t = 544 -> 1024 instead of 544), which costs idle lanes on 26 row
lengths below 8192 that were already correct, including 1536 and 3584.

GGML_OP_NORM is affected as well as GGML_OP_RMS_NORM - both dispatch
through ggml_metal_op_norm.

No mainstream LLM hidden size hits this: ne00_t is ne00/4 on the
vectorized path, so 4096, 8192, 2048 and friends all give a multiple of
32. It is reachable from other norm shapes, e.g. 320-channel norms.

Add NORM and RMS_NORM cases for ne0 = 33, 132 and 260 across the
existing eps values. 33 exercises the scalar path and 132/260 the
vectorized one, since only those divide by 4.

Before, on M3 Pro:

    test-backend-ops test -b MTL0 -o NORM        25/50
    test-backend-ops test -b MTL0 -o RMS_NORM    26/51

After:

    test-backend-ops test -b MTL0 -o NORM        50/50
    test-backend-ops test -b MTL0 -o RMS_NORM    51/51
    test-backend-ops test -b MTL0                13943/13943
2026-08-07 21:09:07 +03:00
Aleksander GrygierandGitHub 23634783c5 ui: Filesystem @mentions for Chat Form (#26715)
* base : @-mention picker foundation - glob search, picker nav, highlight

* feat : @-mention file/folder picker and mention badges in message bubbles

* fix: Imports

* feat : wire the @-mention picker into the chat form

* fix: Bound the glob-search result cache key and prune stale entries
2026-08-07 18:45:54 +02:00
Xuan-Son NguyenandGitHub 4cb22cd537 mtmd: fix longest_edge ignoring min/max pixels (#26638)
* mtmd: fix longest_edge ignoring min/max pixels

* nits
2026-08-07 18:05:15 +02:00
Georgi Gerganov 4cf5cab65d sync : ggml 2026-08-07 17:11:25 +03:00
Georgi Gerganov 933f46f3cb ggml : bump version to 0.19.0 (ggml/1581) 2026-08-07 17:11:25 +03:00
Daniel BeveniusandGitHub 9ba73fd1f5 server : clarify comment in eval_llama_cmpl_schema [no ci] [no release] (#26720) 2026-08-07 15:39:33 +02:00
Emanuil RusevandGitHub f4f7758cae webui: load the model selected via ?model= when ?load=true (#26707)
* webui: load the model selected via ?model=

Opening the WebUI with ?model= selects the model but doesn't load it. The load only starts when you send your first message, so you wait for it then.

This loads it as soon as the page opens, while you're still typing your prompt. It's what the model dropdown already does, and it isn't awaited, so the UI still works while the model loads.

This is the path the Llama macOS app uses to open the WebUI, so it's a common way in.

* webui: gate the load behind ?load=true

Loading on landing is opt-in, so a plain ?model= link behaves as before and doesn't allocate memory on its own.

* webui: name the chat URL params

Collects the query params the chat routes read into a URL_PARAMS constant, instead of repeating the literals across three files. NEW_CHAT_PARAM folds into it.
2026-08-07 15:31:40 +02:00
Niklas WenzelandGitHub 34e9ee57f5 ui: set npm min-release-age to protect against supply-chain attacks (#26711)
* ui: set npm `min-release-age` to protect against supply-chain attacks

* ui: bump to 7 days
2026-08-07 14:53:51 +02:00
Xuan-Son NguyenandGitHub dff15d4ac9 server: (router) add LRU scheduler (#26572)
* add lru_sched

* handle coalescing (req leaves waiting queue)

* add tests

* fix stream case

* address review comments
2026-08-07 14:46:53 +02:00
Xuan-Son NguyenandGitHub e1470ee6a2 server: (router) do not evict busy models (#26567) 2026-08-07 14:39:59 +02:00
PascalandGitHub 217df17ac3 mtmd: stop feeding the text stream again during Qwen3-TTS generation (#26706)
The reference implementation has two mutually exclusive prompt layouts.
In non streaming mode the prefill carries the whole utterance text plus
tts_eos summed with codec_pad, and the trailing text hidden collapses to
a single tts_pad row. In streaming mode the prefill carries only the
first text token and the trailing rows stream the rest of the text
followed by tts_eos.

The pipeline built the non streaming prefill but the streaming overlay,
so the talker saw the utterance a second time during generation and read
it twice before emitting codec_eos.

The overlay is now the single tts_pad row that matches the prefill.
2026-08-07 13:32:52 +02:00
Kilian HuandGitHub cb26014d96 ggml : add aarch64 HWCAP fallbacks and fix fp16 variant detection (#25554)
* ggml : add fallback definitions for missing aarch64 HWCAP bits

* ggml : require HWCAP_ASIMDHP for the aarch64 fp16 cpu variants

Also rename has_fp16_va to has_fp16, the field gates the whole FEAT_FP16
extension, scalar and vector half-precision arithmetic together.
2026-08-07 14:07:10 +03:00
PascalandGitHub 82bb48500a ui: read model modalities from the router model list (#26709)
* ui: read model modalities from the router model list

The router advertises input modalities for every model, loaded or not.
Reading them at list build time lets the UI accept image and audio
uploads for a model selected through ?model=, which has no /props yet.

* enum
2026-08-07 12:07:58 +02:00
Masato NakasakaandGitHub 42e98813e4 Mitigate crashing issue on Windows MSYS2 UCRT64 environment (GCC 16.1.0) (#26555) 2026-08-07 11:17:16 +02:00
Chris LeeandGitHub fc3f10b389 sycl: fix UE4M3 parsing (#25608)
The NVFP4 quantization format stores a scaling factor for every group of
16 weights, packed into a single UE4M3 byte.

The SYCL GPU code was converting these scale values using the E4M3 path,
but that's *signed*, and these are unsigned values.
2026-08-07 08:28:53 +03:00
TitaniumtownandGitHub 6b5c2efb4e sycl: *glu flat path (#26354)
* tests: add SWIGLU perf cases

perf mode had no GLU coverage. Adds SWIGLU at 17408 columns, 512 and
2048 tokens, f16 and f32, with the operands both fused and split.

* sycl: consolidate fused-GLU kernels

They differed only in which op_* they called, so take the op as an argument and share a common launcher.
Their block sizes were all 256, so launch geometry is unchanged;
SYCL_GELU_BLOCK_SIZE and SYCL_SILU_BLOCK_SIZE lose their last users so are dropped.

* sycl: contiguous fast path for the fused GLU ops

o0 == n and o1 == n collapse the de-interleave index math to the
identity, so dispatch a flat kernel in that case. It fires for
ggml_glu_split with packed operands; a fused [gate|up] tensor keeps the
strided path. test-backend-ops perf -o SWIGLU on an Arc Pro B70: split
+14% f16 and +4% f32, fused unchanged.
2026-08-07 08:24:40 +03:00
Neo ZhangandGitHub 31558dbb76 sycl : Support DSv4 OPs: LIGHTNING_INDEXER,DSV4_HC_COMB,DSV4_HC_POST,DSV4_HC_PRE (#26568)
* support DSv4 OPs: LIGHTNING_INDEXER,DSV4_HC_COMB,DSV4_HC_POST,DSV4_HC_PREwq

* update ops.md

* fix format issue
2026-08-07 08:22:23 +03:00
Neo ZhangandGitHub c1f4109898 sycl : update guide Q&A and script for device setting (#26442) 2026-08-07 08:18:47 +03:00
Neo ZhangandGitHub eef5f3e343 sycl : fix error Error OP FLASH_ATTN_EXT on arc770 (#26441) 2026-08-07 08:17:56 +03:00
Neo ZhangandGitHub c074cb3f76 sycl : enhance OP set_rows to support all missed data types (#26515)
* support fp16 to fp16/fp32

* support all missed data types in set_rows

* refactor the code to support all data types
2026-08-07 07:52:52 +03:00
David FriehsandGitHub 5b87ed30f8 cuda: fix warnings for unused variable/function (#26688) 2026-08-07 07:51:56 +03:00
Niklas WenzelandGitHub d8d9887228 ci: abort if build requirements are missing (#26368)
1. Abort CI if build requirements are missing.
2. Add check to make sure Git LFS has been configured.
3. Add trailing newlines to log messages.
2026-08-07 07:50:48 +03:00
JamePengandGitHub e40bf88642 metal : avoid threadgroup matrix array instantiation in kernel_lightning_indexer (#26646)
- In MSL, declaring an array of matrix types like `threadgroup half4x4` causes
a 'no matching constructor' compilation error because MSL matrix types do not
have zero-argument default constructors and threadgroup variables cannot have
initializers.

- Fix this by declaring a POD `threadgroup half` array instead and casting
to `threadgroup half4x4 *` for matrix indexing.

Signed-off-by: JamePeng <jame_peng@sina.com>
2026-08-07 07:49:14 +03:00
Xuan-Son NguyenandGitHub 15586e2d71 mtmd: add chunk save/load function (#26645)
* mtmd: add chunk save/load function

* nits

* add tests

* rn _MAX --> _COUNT
2026-08-06 19:46:40 +02:00
Xuan-Son NguyenandGitHub 6a32c29a74 server: fix empty response for /cors-proxy (#26656) 2026-08-06 15:07:22 +02:00
Sigbjørn SkjæretandGitHub eb5667a169 convert : fix DeepseekV4 rope parameters with transformers 5.x (#26673) 2026-08-06 16:06:52 +03:00
118 changed files with 28270 additions and 1815 deletions
+22 -9
View File
@@ -643,39 +643,52 @@ function gg_sum_rerank_tiny {
function gg_check_build_requirements {
if ! command -v git &> /dev/null; then
gg_printf 'git not found, please install'
gg_printf 'git not found, please install\n'
exit 1
fi
if ! command -v git-lfs &> /dev/null; then
gg_printf 'git-lfs not found, please install'
gg_printf 'git-lfs not found, please install\n'
exit 1
fi
if ! git config --get filter.lfs.clean &> /dev/null; then
gg_printf 'git-lfs not initialized, please run `git lfs install`\n'
exit 1
fi
if ! command -v wget &> /dev/null; then
gg_printf 'wget not found, please install'
gg_printf 'wget not found, please install\n'
exit 1
fi
if ! command -v python3 &> /dev/null; then
gg_printf 'python3 not found, please install'
gg_printf 'python3 not found, please install\n'
exit 1
fi
if ! command -v pip3 &> /dev/null; then
gg_printf 'pip3 not found, please install'
gg_printf 'pip3 not found, please install\n'
exit 1
fi
if ! python3 -m ensurepip --help &> /dev/null; then
gg_printf 'ensurepip not found, please install python3-venv package'
gg_printf 'ensurepip not found, please install python3-venv package\n'
exit 1
fi
if ! command -v cmake &> /dev/null; then
gg_printf 'cmake not found, please install'
gg_printf 'cmake not found, please install\n'
exit 1
fi
if ! command -v ccache &> /dev/null; then
gg_printf 'ccache not found, please consider installing for faster builds'
gg_printf 'ccache not found, please consider installing for faster builds\n'
fi
if ! command -v ctest &> /dev/null; then
gg_printf 'ctest not found, please install'
gg_printf 'ctest not found, please install\n'
exit 1
fi
}
+7
View File
@@ -533,6 +533,13 @@ class DeepseekV4Model(TextModel):
for key, value in raw_hparams.items():
self.hparams.setdefault(key, value)
# workaround for special rope_parameters (main/compress) in transformers 5.x
if self.rope_parameters.get("full_attention", self.rope_parameters).get("rope_type") is None:
if (rope_scaling := raw_hparams.get("rope_scaling")) is not None:
if "rope_type" not in rope_scaling and (rope_type := rope_scaling.get("type")) is not None:
rope_scaling["rope_type"] = rope_type
self.rope_parameters.update(**rope_scaling)
self.block_count = self.hparams["num_hidden_layers"]
if self.mtp_only:
self.block_count += self.hparams.get("num_nextn_predict_layers", 0)
+42
View File
@@ -449,6 +449,8 @@ Or
use 1 SYCL GPUs: [0] with Max compute units:512
```
User can use the device management in [docs/multi-gpu.md](https://github.com/ggml-org/llama.cpp/blob/master/docs/multi-gpu.md), like parameter `--device SYCL0,SYCL1` to assign one or more devices.
## Windows
### Install GPU driver
@@ -763,6 +765,7 @@ Or
use 1 SYCL GPUs: [0] with Max compute units:512
```
User can use the device management in [docs/multi-gpu.md](https://github.com/ggml-org/llama.cpp/blob/master/docs/multi-gpu.md), like parameter `--device SYCL0,SYCL1` to assign one or more devices.
## Environment Variable
@@ -895,6 +898,45 @@ Pass these via `CXXFLAGS` or add a one-off `#define` to enable a flag on the spo
set UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=1
```
- When I set `SYCL_CACHE_PERSISTENT=1` in running time, I meet crash.
`SYCL_CACHE_PERSISTENT=1` is not recommended by llama.cpp SYCL backend.
When cache is enabled, SYCL runtime will try to cache and reuse JIT-compiled binaries.
We find some AI will tell user this cmd to speed up SYCL backend. It only speeds up the startup to skip the JIT process, instead of running speed.
It will bring negative impact when the SYCL binary file is changed frequently in your running environment. The new & old codes mix will lead to crash.
Compare to the benefit, it has brought more failed cases.
If you are not familiar with the SYCL compiler principle of JIT and AOT, please don't use it.
To restore, you need to remove the local cache: `~/.cache/libsycl_cache/` and execute `unset SYCL_CACHE_PERSISTENT` in running time.
- How to use iGPU and dGPU in same time?
1. Detect the devices in your running time.
```
source /opt/intel/oneapi/setvars.sh
./build/bin/llama-server --list-devices
or
./build/bin/llama-cli --list-devices
./build/bin/llama-bench --list-devices
./build/bin/llama-completion --list-devices
Available devices:
SYCL0: Intel(R) Arc(TM) A770 Graphics (15473 MiB, 15473 MiB free)
SYCL1: Intel(R) UHD Graphics 770 (59675 MiB, 44986 MiB free)
```
The dGPU will be in the head of this list and iGPU will be the end.
If not all GPUs are listed, please check the env var: ONEAPI_DEVICE_SELECTOR and unset it.
2. Set the iGPU and dGPU
Set the iGPU and dGPU by `./build/bin/llama-server --device SYCL0,SYCL1,SYCLxxx`.
### **GitHub contribution**:
Please add the `[SYCL]` prefix/tag in issues/PRs titles to help the SYCL contributors to check/address them without delay.
+6 -6
View File
@@ -15,7 +15,7 @@ Legend:
| Operation | BLAS | CANN | CPU | CUDA | ET | MTL | OpenCL | SYCL | Vulkan | WebGPU | ZenDNN | zDNN |
|-----------|------|------|------|------|------|------|------|------|------|------|------|------|
| ABS | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ACC | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | 🟡 | ✅ | ❌ | ❌ | ❌ |
| ACC | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | | ✅ | ❌ | ❌ | ❌ |
| ADD | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| ADD1 | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ |
| ADD_ID | ❌ | ❌ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -41,9 +41,9 @@ Legend:
| DIAG | ❌ | ❌ | ✅ | ✅ | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| DIAG_MASK_INF | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ❌ |
| DIV | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
| DSV4_HC_COMB | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DSV4_HC_POST | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DSV4_HC_PRE | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DSV4_HC_COMB | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DSV4_HC_POST | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DSV4_HC_PRE | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| DUP | ❌ | ✅ | ✅ | 🟡 | ❌ | 🟡 | 🟡 | ✅ | ✅ | ❌ | ❌ | ❌ |
| ELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| EXP | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -59,7 +59,7 @@ Legend:
| GELU | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_ERF | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GELU_QUICK | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | 🟡 | ✅ | ✅ | ✅ | ❌ | ❌ |
| GET_ROWS | ❌ | 🟡 | ✅ | 🟡 | 🟡 | 🟡 | 🟡 | | ✅ | 🟡 | ❌ | ❌ |
| GET_ROWS | ❌ | 🟡 | ✅ | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | ✅ | 🟡 | ❌ | ❌ |
| GET_ROWS_BACK | ❌ | ❌ | 🟡 | 🟡 | ❌ | ❌ | ❌ | ❌ | 🟡 | ❌ | ❌ | ❌ |
| GROUP_NORM | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| HARDSIGMOID | ❌ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
@@ -68,7 +68,7 @@ Legend:
| IM2COL_3D | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| L2_NORM | ❌ | ✅ | ✅ | ✅ | 🟡 | ✅ | ❌ | ✅ | ✅ | 🟡 | ❌ | ❌ |
| LEAKY_RELU | ❌ | ✅ | ✅ | ✅ | ❌ | 🟡 | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
| LIGHTNING_INDEXER | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| LIGHTNING_INDEXER | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | ❌ | | ❌ | ❌ | ❌ | ❌ |
| LOG | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
| MEAN | ❌ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ |
| MUL | ❌ | ✅ | ✅ | ✅ | 🟡 | 🟡 | ✅ | ✅ | ✅ | ✅ | ❌ | ❌ |
+22870 -671
View File
File diff suppressed because it is too large Load Diff
+14 -7
View File
@@ -12,6 +12,7 @@ This script processes files with specified options.
Options:
-h, --help Display this help message and exit.
-d, --device <value> Set SYCL devices (default: SYCL0).
-c, --context <value> Set context length. Bigger need more memory.
-p, --promote <value> Prompt to start generation with.
-m, --model <value> Full model file path.
@@ -41,10 +42,16 @@ MODEL_FILE=../models/Qwen3.5-4B-Q4_0.gguf
NGL=99
CONTEXT=4096
GGML_SYCL_DEVICE=-1
SYCL_DEVICES="SYCL0"
SPLIT_MODE=layer
LOG_VERBOSE=3
while [[ $# -gt 0 ]]; do
case "$1" in
-d|--device)
SYCL_DEVICES="$2"
shift
shift
;;
-c|--context)
CONTEXT=$2
# Shift twice to consume both the option flag and its value
@@ -95,8 +102,6 @@ while [[ $# -gt 0 ]]; do
esac
done
source /opt/intel/oneapi/setvars.sh
#export GGML_SYCL_DEBUG=1
@@ -107,17 +112,19 @@ source /opt/intel/oneapi/setvars.sh
export UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=1
echo "UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=${UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS}"
echo "ONEAPI_DEVICE_SELECTOR=${ONEAPI_DEVICE_SELECTOR}"
if [ $GGML_SYCL_DEVICE -ne -1 ]; then
echo "Use $GGML_SYCL_DEVICE as main GPU"
#use signle GPU only
GPUS_SETTING="-mg $GGML_SYCL_DEVICE -sm ${SPLIT_MODE}"
echo "ONEAPI_DEVICE_SELECTOR=${ONEAPI_DEVICE_SELECTOR}"
else
echo "Use all Intel GPUs, including iGPU & dGPU"
echo "Use Intel GPUs: ${SYCL_DEVICES}"
GPUS_SETTING="-sm ${SPLIT_MODE}"
fi
fi
echo "run cmd: ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --mmap --host 0.0.0.0 --port 8000"
ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --mmap --host 0.0.0.0 --port 8000
echo "run cmd: ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --device ${SYCL_DEVICES} --mmap --host 0.0.0.0 --port 8000"
ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --device ${SYCL_DEVICES} --mmap --host 0.0.0.0 --port 8000
+12 -4
View File
@@ -12,6 +12,7 @@ This script processes files with specified options.
Options:
-h, --help Display this help message and exit.
-d, --device <value> Set SYCL devices (default: SYCL0).
-c, --context <value> Set context length. Bigger need more memory.
-p, --promote <value> Prompt to start generation with.
-m, --model <value> Full model file path.
@@ -42,10 +43,16 @@ MODEL_FILE=../models/llama-2-7b.Q4_0.gguf
NGL=99
CONTEXT=4096
GGML_SYCL_DEVICE=-1
SYCL_DEVICES="SYCL0"
SPLIT_MODE=layer
LOG_VERBOSE=3
while [[ $# -gt 0 ]]; do
case "$1" in
-d|--device)
SYCL_DEVICES="$2"
shift
shift
;;
-c|--context)
CONTEXT=$2
# Shift twice to consume both the option flag and its value
@@ -115,16 +122,17 @@ source /opt/intel/oneapi/setvars.sh
export UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=1
echo "UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=${UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS}"
echo "ONEAPI_DEVICE_SELECTOR=${ONEAPI_DEVICE_SELECTOR}"
if [ $GGML_SYCL_DEVICE -ne -1 ]; then
echo "Use $GGML_SYCL_DEVICE as main GPU"
#use signle GPU only
GPUS_SETTING="-mg $GGML_SYCL_DEVICE -sm ${SPLIT_MODE}"
echo "ONEAPI_DEVICE_SELECTOR=${ONEAPI_DEVICE_SELECTOR}"
else
echo "Use all Intel GPUs, including iGPU & dGPU"
echo "Use Intel GPUs: ${SYCL_DEVICES}"
GPUS_SETTING="-sm ${SPLIT_MODE}"
fi
echo "run cmd: ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -no-cnv -p "${INPUT_PROMPT}" -n 200 -e -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --mmap "
ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -no-cnv -p "${INPUT_PROMPT}" -n 200 -e -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --mmap
echo "run cmd: ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -no-cnv -p "${INPUT_PROMPT}" -n 200 -e -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --device ${SYCL_DEVICES} --mmap "
ZES_ENABLE_SYSMAN=1 ${BIN_FILE} -m ${MODEL_FILE} -no-cnv -p "${INPUT_PROMPT}" -n 200 -e -ngl ${NGL} -s ${SEED} -c ${CONTEXT} ${GPUS_SETTING} -lv ${LOG_VERBOSE} --device ${SYCL_DEVICES} --mmap
+23 -5
View File
@@ -13,6 +13,7 @@ set "MODEL_FILE=..\models\Qwen3.5-4B-Q4_0.gguf"
set "NGL=99"
set "CONTEXT=4096"
set "GGML_SYCL_DEVICE=-1"
set "SYCL_DEVICES=SYCL0"
set "SPLIT_MODE=layer"
set "LOG_VERBOSE=3"
@@ -36,6 +37,21 @@ if /I "%~1"=="--context" (
goto parse_args
)
if /I "%~1"=="-d" (
if "%~2"=="" goto missing_value
set "SYCL_DEVICES=%~2"
shift
shift
goto parse_args
)
if /I "%~1"=="--device" (
if "%~2"=="" goto missing_value
set "SYCL_DEVICES=%~2"
shift
shift
goto parse_args
)
if /I "%~1"=="-m" (
if "%~2"=="" goto missing_value
set "MODEL_FILE=%~2"
@@ -130,6 +146,7 @@ echo This script processes files with specified options.
echo.
echo Options:
echo -h, --help Display this help message and exit.
echo -d, --device ^<value^> Set SYCL devices (default: SYCL0).
echo -c, --context ^<value^> Set context length. Bigger need more memory.
echo -m, --model ^<value^> Full model file path.
echo -mg,--main-gpu ^<value^> Set main GPU ID (0 - n) for single GPU mode.
@@ -160,19 +177,20 @@ REM Support malloc device memory more than 4GB.
set "UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=1"
echo UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=%UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS%
echo ONEAPI_DEVICE_SELECTOR=%ONEAPI_DEVICE_SELECTOR%
if not "%GGML_SYCL_DEVICE%"=="-1" (
echo Use %GGML_SYCL_DEVICE% as main GPU
REM Use single GPU only.
set "GPUS_SETTING=-mg %GGML_SYCL_DEVICE% -sm %SPLIT_MODE%"
echo ONEAPI_DEVICE_SELECTOR=%ONEAPI_DEVICE_SELECTOR%
) else (
echo Use all Intel GPUs, including iGPU ^& dGPU
) else (
echo Use Intel GPUs: %SYCL_DEVICES%
set "GPUS_SETTING=-sm %SPLIT_MODE%"
)
echo run cmd: ZES_ENABLE_SYSMAN=1 %BIN_FILE% -m "%MODEL_FILE%" -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --mmap --host 0.0.0.0 --port 8000
echo run cmd: ZES_ENABLE_SYSMAN=1 %BIN_FILE% -m "%MODEL_FILE%" -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --device %SYCL_DEVICES% --mmap --host 0.0.0.0 --port 8000
set "ZES_ENABLE_SYSMAN=1"
%BIN_FILE% -m "%MODEL_FILE%" -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --mmap --host 0.0.0.0 --port 8000
%BIN_FILE% -m "%MODEL_FILE%" -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --device "%SYCL_DEVICES%" --mmap --host 0.0.0.0 --port 8000
endlocal
+24 -5
View File
@@ -19,6 +19,7 @@ set "MODEL_FILE=..\models\llama-2-7b.Q4_0.gguf"
set "NGL=99"
set "CONTEXT=4096"
set "GGML_SYCL_DEVICE=-1"
set "SYCL_DEVICES=SYCL0"
set "SPLIT_MODE=layer"
set "LOG_VERBOSE=3"
@@ -42,6 +43,21 @@ if /I "%~1"=="--context" (
goto parse_args
)
if /I "%~1"=="-d" (
if "%~2"=="" goto missing_value
set "SYCL_DEVICES=%~2"
shift
shift
goto parse_args
)
if /I "%~1"=="--device" (
if "%~2"=="" goto missing_value
set "SYCL_DEVICES=%~2"
shift
shift
goto parse_args
)
if /I "%~1"=="-p" (
if "%~2"=="" goto missing_value
set "INPUT_PROMPT=%~2"
@@ -151,6 +167,7 @@ echo This script processes files with specified options.
echo.
echo Options:
echo -h, --help Display this help message and exit.
echo -d, --device ^<value^> Set SYCL devices (default: SYCL0).
echo -c, --context ^<value^> Set context length. Bigger need more memory.
echo -p, --promote ^<value^> Prompt to start generation with.
echo -m, --model ^<value^> Full model file path.
@@ -182,19 +199,21 @@ REM Support malloc device memory more than 4GB.
set "UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=1"
echo UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS=%UR_L0_ENABLE_RELAXED_ALLOCATION_LIMITS%
echo ONEAPI_DEVICE_SELECTOR=%ONEAPI_DEVICE_SELECTOR%
if not "%GGML_SYCL_DEVICE%"=="-1" (
echo Use %GGML_SYCL_DEVICE% as main GPU
REM Use single GPU only.
set "GPUS_SETTING=-mg %GGML_SYCL_DEVICE% -sm %SPLIT_MODE%"
echo ONEAPI_DEVICE_SELECTOR=%ONEAPI_DEVICE_SELECTOR%
) else (
echo Use all Intel GPUs, including iGPU ^& dGPU
)
else (
echo Use Intel GPUs: %SYCL_DEVICES%
set "GPUS_SETTING=-sm %SPLIT_MODE%"
)
echo run cmd: ZES_ENABLE_SYSMAN=1 %BIN_FILE% -m %MODEL_FILE% -no-cnv -p "%INPUT_PROMPT%" -n 200 -e -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --mmap
echo run cmd: ZES_ENABLE_SYSMAN=1 %BIN_FILE% -m %MODEL_FILE% -no-cnv -p "%INPUT_PROMPT%" -n 200 -e -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --device %SYCL_DEVICES% --mmap
set "ZES_ENABLE_SYSMAN=1"
%BIN_FILE% -m "%MODEL_FILE%" -no-cnv -p "%INPUT_PROMPT%" -n 200 -e -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --mmap
%BIN_FILE% -m "%MODEL_FILE%" -no-cnv -p "%INPUT_PROMPT%" -n 200 -e -ngl %NGL% -s %SEED% -c %CONTEXT% %GPUS_SETTING% -lv %LOG_VERBOSE% --device "%SYCL_DEVICES%" --mmap
endlocal
+2 -2
View File
@@ -4,8 +4,8 @@ project("ggml" C CXX ASM)
### GGML Version
set(GGML_VERSION_MAJOR 0)
set(GGML_VERSION_MINOR 18)
set(GGML_VERSION_PATCH 1)
set(GGML_VERSION_MINOR 19)
set(GGML_VERSION_PATCH 0)
set(GGML_VERSION_BASE "${GGML_VERSION_MAJOR}.${GGML_VERSION_MINOR}.${GGML_VERSION_PATCH}")
list(APPEND CMAKE_MODULE_PATH "${CMAKE_CURRENT_SOURCE_DIR}/cmake/")
+19 -3
View File
@@ -8,6 +8,22 @@
#include <sys/sysctl.h>
#endif
#if !defined(HWCAP_FPHP)
#define HWCAP_FPHP (1 << 9)
#endif
#if !defined(HWCAP_ASIMDHP)
#define HWCAP_ASIMDHP (1 << 10)
#endif
#if !defined(HWCAP_ASIMDDP)
#define HWCAP_ASIMDDP (1 << 20)
#endif
#if !defined(HWCAP_SVE)
#define HWCAP_SVE (1 << 22)
#endif
#if !defined(HWCAP2_SVE2)
#define HWCAP2_SVE2 (1 << 1)
#endif
@@ -23,7 +39,7 @@
struct aarch64_features {
// has_neon not needed, aarch64 has NEON guaranteed
bool has_dotprod = false;
bool has_fp16_va = false;
bool has_fp16 = false;
bool has_sve = false;
bool has_sve2 = false;
bool has_i8mm = false;
@@ -36,7 +52,7 @@ struct aarch64_features {
uint32_t hwcap2 = getauxval(AT_HWCAP2);
has_dotprod = !!(hwcap & HWCAP_ASIMDDP);
has_fp16_va = !!(hwcap & HWCAP_FPHP);
has_fp16 = !!(hwcap & HWCAP_FPHP) && !!(hwcap & HWCAP_ASIMDHP);
has_sve = !!(hwcap & HWCAP_SVE);
has_sve2 = !!(hwcap2 & HWCAP2_SVE2);
has_i8mm = !!(hwcap2 & HWCAP2_I8MM);
@@ -75,7 +91,7 @@ static int ggml_backend_cpu_aarch64_score() {
score += 1<<1;
#endif
#ifdef GGML_USE_FP16_VECTOR_ARITHMETIC
if (!af.has_fp16_va) { return 0; }
if (!af.has_fp16) { return 0; }
score += 1<<2;
#endif
#ifdef GGML_USE_SVE
+1
View File
@@ -5209,6 +5209,7 @@ static bool ggml_backend_cuda_device_offload_op(ggml_backend_dev_t dev, const gg
static ggml_backend_event_t ggml_backend_cuda_device_event_new(ggml_backend_dev_t dev) {
#ifdef GGML_CUDA_NO_PEER_COPY
GGML_UNUSED(dev);
return nullptr;
#else
ggml_backend_cuda_device_context * dev_ctx = (ggml_backend_cuda_device_context *)dev->context;
+1 -1
View File
@@ -8,7 +8,6 @@ struct __builtin_align__(32) float8 {
float x; float y; float z; float w;
float p; float q; float r; float s;
};
#endif
#if CUDART_VERSION >= 12080
static __device__ __forceinline__ float nvfp4_native_scale_error(
@@ -49,6 +48,7 @@ static __device__ __forceinline__ float nvfp4_native_scale_error(
return err;
}
#endif // CUDART_VERSION >= 12080
#endif // defined(BLACKWELL_MMA_AVAILABLE)
__launch_bounds__(CUDA_QUANTIZE_BLOCK_SIZE, 1)
static __global__ void quantize_q8_1(
+1 -1
View File
@@ -3816,7 +3816,7 @@ int ggml_metal_op_norm(ggml_metal_op_t ctx, int idx) {
}
nth = std::min(nth, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
nth = std::min(nth, args.ne00_t);
nth = std::min(nth, (args.ne00_t + 31)/32*32);
const size_t smem = pipeline.smem;
+2 -2
View File
@@ -11328,8 +11328,8 @@ kernel void kernel_lightning_indexer(
const int i_kv_0 = tgpig.x*NK; // first key of this threadgroup
const int i_kv = i_kv_0 + sgitg*NKPSG; // first key of this simdgroup
threadgroup half4x4 sk4x4[NK*DK16];
threadgroup half * sk = (threadgroup half *) sk4x4;
threadgroup half sk[NK * DK16 * 16];
threadgroup half4x4 * sk4x4 = (threadgroup half4x4 *) sk;
for (short i = tiitg; i < NK*DK16; i += NTG) {
const short ik = i/DK16;
+14 -3
View File
@@ -1022,9 +1022,20 @@ static T block_reduce(T val, T * shared_vals, int block_size_template) {
}
static __dpct_inline__ float ggml_sycl_ue4m3_to_fp32(uint8_t x) {
const uint32_t bits = x * (x != 0x7F && x != 0xFF);
const __nv_fp8_e4m3 xf = *reinterpret_cast<const __nv_fp8_e4m3 *>(&bits);
return static_cast<float>(xf) / 2;
// UE4M3 is unsigned: 4 exp bits (bias 7), 3 mantissa bits, no sign, no NaN.
// exp == 0xF is a valid exponent (256-448 range), not NaN.
if (x == 0 || x == 0x7F) {
return 0.0f;
}
const int exp = (x >> 3) & 0xF;
const int man = x & 0x7;
float raw;
if (exp == 0) {
raw = man * (1.0f / 8.0f) * sycl::pow(2.0f, -6.0f);
} else {
raw = (1.0f + man / 8.0f) * sycl::pow(2.0f, (float) exp - 7.0f);
}
return raw * 0.5f;
}
#endif // GGML_SYCL_COMMON_HPP
+280
View File
@@ -0,0 +1,280 @@
#include "ggml-impl.h"
#include "dsv4-hc.hpp"
#include <cmath>
static constexpr int DSV4_HC = 4;
static void dsv4_hc_pre_f32_sycl(
const float * x, const float * weights, float * dst,
int64_t n_embd, int64_t hc, int64_t n_tokens,
int64_t sx0, int64_t sx1, int64_t sx2,
int64_t sw0, int64_t sw1,
int64_t sd0, int64_t sd1,
queue_ptr stream) {
const int64_t nr = n_embd * n_tokens;
const int64_t block_size = 256;
const int64_t num_blocks = (nr + block_size - 1) / block_size;
stream->parallel_for(
sycl::nd_range<1>(sycl::range<1>(num_blocks * block_size), sycl::range<1>(block_size)),
[=](sycl::nd_item<1> item) {
const int64_t ir = item.get_global_id(0);
if (ir >= nr) {
return;
}
const int64_t i0 = ir % n_embd;
const int64_t it = ir / n_embd;
float sum = x[i0*sx0 + it*sx2] * weights[it*sw1];
for (int64_t ih = 1; ih < hc; ++ih) {
const float xv = x[i0*sx0 + ih*sx1 + it*sx2];
const float wv = weights[ih*sw0 + it*sw1];
sum += xv * wv;
}
dst[i0*sd0 + it*sd1] = sum;
});
}
static void dsv4_hc_comb_norm_cols(float * comb, float eps) {
for (int idst = 0; idst < DSV4_HC; ++idst) {
float sum = eps;
for (int isrc = 0; isrc < DSV4_HC; ++isrc) {
sum += comb[idst + DSV4_HC*isrc];
}
const float inv_sum = 1.0f / sum;
for (int isrc = 0; isrc < DSV4_HC; ++isrc) {
comb[idst + DSV4_HC*isrc] *= inv_sum;
}
}
}
static void dsv4_hc_comb_norm_rows(float * comb, float eps) {
for (int isrc = 0; isrc < DSV4_HC; ++isrc) {
float sum = eps;
for (int idst = 0; idst < DSV4_HC; ++idst) {
sum += comb[idst + DSV4_HC*isrc];
}
const float inv_sum = 1.0f / sum;
for (int idst = 0; idst < DSV4_HC; ++idst) {
comb[idst + DSV4_HC*isrc] *= inv_sum;
}
}
}
static void dsv4_hc_comb_f32_sycl(
const float * mixes,
const float * scale,
const float * base,
float * dst,
int64_t n_tokens,
int64_t sm0,
int64_t sm1,
int64_t ss0,
int64_t sb0,
int64_t sd0,
int64_t sd1,
int64_t sd2,
float eps,
int32_t n_iter,
queue_ptr stream) {
constexpr int comb_offset = 2*DSV4_HC;
const int64_t block_size = 256;
const int64_t num_blocks = (n_tokens + block_size - 1) / block_size;
stream->parallel_for(
sycl::nd_range<1>(sycl::range<1>(num_blocks * block_size), sycl::range<1>(block_size)),
[=](sycl::nd_item<1> item_ct1) {
const int64_t it = item_ct1.get_global_id(0);
if (it >= n_tokens) {
return;
}
const float scale_comb = scale[2*ss0];
float comb[DSV4_HC*DSV4_HC];
for (int isrc = 0; isrc < DSV4_HC; ++isrc) {
float max = -INFINITY;
for (int idst = 0; idst < DSV4_HC; ++idst) {
const int idx = idst + DSV4_HC*isrc;
const float v = mixes[(comb_offset + idx)*sm0 + it*sm1] * scale_comb + base[(comb_offset + idx)*sb0];
comb[idx] = v;
max = fmaxf(max, v);
}
float sum = 0.0f;
for (int idst = 0; idst < DSV4_HC; ++idst) {
const int idx = idst + DSV4_HC*isrc;
const float v = expf(comb[idx] - max);
comb[idx] = v;
sum += v;
}
const float inv_sum = 1.0f / sum;
for (int idst = 0; idst < DSV4_HC; ++idst) {
const int idx = idst + DSV4_HC*isrc;
comb[idx] = comb[idx] * inv_sum + eps;
}
}
dsv4_hc_comb_norm_cols(comb, eps);
for (int32_t i = 1; i < n_iter; ++i) {
dsv4_hc_comb_norm_rows(comb, eps);
dsv4_hc_comb_norm_cols(comb, eps);
}
for (int isrc = 0; isrc < DSV4_HC; ++isrc) {
for (int idst = 0; idst < DSV4_HC; ++idst) {
const int idx = idst + DSV4_HC*isrc;
dst[idst*sd0 + isrc*sd1 + it*sd2] = comb[idx];
}
}
});
}
static void dsv4_hc_post_f32_sycl(
const float * x, const float * residual, const float * post, const float * comb, float * dst,
int64_t n_embd, int64_t hc, int64_t n_tokens,
int64_t sx0, int64_t sx1,
int64_t sr0, int64_t sr1, int64_t sr2,
int64_t sp0, int64_t sp1,
int64_t sc0, int64_t sc1, int64_t sc2,
int64_t sd0, int64_t sd1, int64_t sd2,
queue_ptr stream) {
const int64_t nr = n_embd * hc * n_tokens;
const int64_t block_size = 256;
const int64_t num_blocks = (nr + block_size - 1) / block_size;
stream->parallel_for(
sycl::nd_range<1>(sycl::range<1>(num_blocks * block_size), sycl::range<1>(block_size)),
[=](sycl::nd_item<1> item) {
const int64_t ir = item.get_global_id(0);
if (ir >= nr) {
return;
}
const int64_t i0 = ir % n_embd;
const int64_t idst = (ir / n_embd) % hc;
const int64_t it = ir / (n_embd * hc);
float sum = x[i0*sx0 + it*sx1] * post[idst*sp0 + it*sp1];
for (int64_t isrc = 0; isrc < hc; ++isrc) {
sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2];
}
dst[i0*sd0 + idst*sd1 + it*sd2] = sum;
});
}
void ggml_sycl_op_dsv4_hc_pre(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/2);
const ggml_tensor * x = dst->src[0];
const ggml_tensor * weights = dst->src[1];
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(weights->type == GGML_TYPE_F32);
GGML_ASSERT(dst->type == GGML_TYPE_F32);
GGML_TENSOR_LOCALS(size_t, nbx, x, nb);
GGML_TENSOR_LOCALS(size_t, nbw, weights, nb);
GGML_TENSOR_LOCALS(size_t, nbd, dst, nb);
const int64_t n_embd = x->ne[0];
const int64_t hc = x->ne[1];
const int64_t n_tokens = x->ne[2];
queue_ptr stream = ctx.stream();
dsv4_hc_pre_f32_sycl(
(const float *) x->data, (const float *) weights->data, (float *) dst->data,
n_embd, hc, n_tokens,
nbx0 / sizeof(float), nbx1 / sizeof(float), nbx2 / sizeof(float),
nbw0 / sizeof(float), nbw1 / sizeof(float),
nbd0 / sizeof(float), nbd1 / sizeof(float),
stream);
}
void ggml_sycl_op_dsv4_hc_comb(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/3);
const ggml_tensor * mixes = dst->src[0];
const ggml_tensor * scale = dst->src[1];
const ggml_tensor * base = dst->src[2];
GGML_ASSERT(mixes->type == GGML_TYPE_F32);
GGML_ASSERT(scale->type == GGML_TYPE_F32);
GGML_ASSERT(base->type == GGML_TYPE_F32);
GGML_ASSERT(dst->type == GGML_TYPE_F32);
constexpr int64_t hc_mix_dim = (2 + DSV4_HC)*DSV4_HC;
GGML_ASSERT(mixes->ne[0] == hc_mix_dim);
GGML_ASSERT(dst->ne[0] == DSV4_HC);
GGML_ASSERT(dst->ne[1] == DSV4_HC);
GGML_ASSERT(dst->ne[2] == mixes->ne[1]);
GGML_ASSERT(scale->ne[0] >= 3);
GGML_ASSERT(base->ne[0] == hc_mix_dim);
GGML_TENSOR_LOCALS(size_t, nbm, mixes, nb);
GGML_TENSOR_LOCALS(size_t, nbs, scale, nb);
GGML_TENSOR_LOCALS(size_t, nbb, base, nb);
GGML_TENSOR_LOCALS(size_t, nbd, dst, nb);
const int64_t n_tokens = mixes->ne[1];
const float eps = ggml_get_op_params_f32(dst, 0);
const int32_t n_iter = ggml_get_op_params_i32(dst, 1);
queue_ptr stream = ctx.stream();
dsv4_hc_comb_f32_sycl(
(const float *) mixes->data, (const float *) scale->data, (const float *) base->data, (float *) dst->data,
n_tokens,
nbm0 / sizeof(float), nbm1 / sizeof(float),
nbs0 / sizeof(float),
nbb0 / sizeof(float),
nbd0 / sizeof(float), nbd1 / sizeof(float), nbd2 / sizeof(float),
eps, n_iter, stream);
}
void ggml_sycl_op_dsv4_hc_post(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/4);
const ggml_tensor * x = dst->src[0];
const ggml_tensor * residual = dst->src[1];
const ggml_tensor * post = dst->src[2];
const ggml_tensor * comb = dst->src[3];
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(residual->type == GGML_TYPE_F32);
GGML_ASSERT(post->type == GGML_TYPE_F32);
GGML_ASSERT(comb->type == GGML_TYPE_F32);
GGML_ASSERT(dst->type == GGML_TYPE_F32);
GGML_TENSOR_LOCALS(size_t, nbx, x, nb);
GGML_TENSOR_LOCALS(size_t, nbr, residual, nb);
GGML_TENSOR_LOCALS(size_t, nbp, post, nb);
GGML_TENSOR_LOCALS(size_t, nbc, comb, nb);
GGML_TENSOR_LOCALS(size_t, nbd, dst, nb);
const int64_t n_embd = x->ne[0];
const int64_t n_tokens = x->ne[1];
const int64_t hc = residual->ne[1];
queue_ptr stream = ctx.stream();
dsv4_hc_post_f32_sycl(
(const float *) x->data, (const float *) residual->data,
(const float *) post->data, (const float *) comb->data, (float *) dst->data,
n_embd, hc, n_tokens,
nbx0 / sizeof(float), nbx1 / sizeof(float),
nbr0 / sizeof(float), nbr1 / sizeof(float), nbr2 / sizeof(float),
nbp0 / sizeof(float), nbp1 / sizeof(float),
nbc0 / sizeof(float), nbc1 / sizeof(float), nbc2 / sizeof(float),
nbd0 / sizeof(float), nbd1 / sizeof(float), nbd2 / sizeof(float),
stream);
}
+10
View File
@@ -0,0 +1,10 @@
#ifndef GGML_SYCL_DSV4_HC_HPP
#define GGML_SYCL_DSV4_HC_HPP
#include "common.hpp"
void ggml_sycl_op_dsv4_hc_pre(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
void ggml_sycl_op_dsv4_hc_comb(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
void ggml_sycl_op_dsv4_hc_post(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
#endif // GGML_SYCL_DSV4_HC_HPP
+65 -93
View File
@@ -420,53 +420,31 @@ static void clamp(const T * x, T * dst, const float min, const float max, const
}
}
template<typename T>
static void gated_op_fused_geglu(const T * x, const T * g, T * dst, const uint64_t k, const sycl::uint3 n_fd, const uint64_t o0, const uint64_t o1, const sycl::nd_item<1> &item_ct1) {
template<typename T, typename F>
static void unary_gated_op_flat_kernel(const T * x, const T * g, T * dst, const uint64_t k, const sycl::nd_item<1> & item_ct1, F func) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
dst[i] = func(x[i]) * g[i];
}
}
template<typename T, typename F>
static void unary_gated_op_generic_kernel(
const T * x,
const T * g,
T * dst,
const uint64_t k,
const sycl::uint3 n_fd,
const uint64_t o0,
const uint64_t o1,
const sycl::nd_item<1> & item_ct1,
F func) {
// rows of n columns at strides o0 and o1: two halves of one fused tensor, or two tensors
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const sycl::uint2 rc = fast_div_modulo((uint32_t) i, n_fd);
const int64_t j0 = rc.x() * o0 + rc.y();
const int64_t j1 = o0 == o1 ? j0 : rc.x() * o1 + rc.y();
dst[i] = op_gelu(x[j0]) * g[j1];
}
}
template<typename T>
static void gated_op_fused_reglu(const T * x, const T * g, T * dst, const uint64_t k, const sycl::uint3 n_fd, const uint64_t o0, const uint64_t o1, const sycl::nd_item<1> &item_ct1) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const sycl::uint2 rc = fast_div_modulo((uint32_t) i, n_fd);
const int64_t j0 = rc.x() * o0 + rc.y();
const int64_t j1 = o0 == o1 ? j0 : rc.x() * o1 + rc.y();
dst[i] = op_relu(x[j0]) * g[j1];
}
}
template<typename T>
static void gated_op_fused_swiglu(const T * x, const T * g, T * dst, const uint64_t k, const sycl::uint3 n_fd, const uint64_t o0, const uint64_t o1, const sycl::nd_item<1> &item_ct1) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const sycl::uint2 rc = fast_div_modulo((uint32_t) i, n_fd);
const int64_t j0 = rc.x() * o0 + rc.y();
const int64_t j1 = o0 == o1 ? j0 : rc.x() * o1 + rc.y();
dst[i] = op_silu(x[j0]) * g[j1];
}
}
template<typename T>
static void gated_op_fused_geglu_erf(const T * x, const T * g, T * dst, const uint64_t k, const sycl::uint3 n_fd, const uint64_t o0, const uint64_t o1, const sycl::nd_item<1> &item_ct1) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const sycl::uint2 rc = fast_div_modulo((uint32_t) i, n_fd);
const int64_t j0 = rc.x() * o0 + rc.y();
const int64_t j1 = o0 == o1 ? j0 : rc.x() * o1 + rc.y();
dst[i] = op_gelu_erf(x[j0]) * g[j1];
}
}
template<typename T>
static void gated_op_fused_geglu_quick(const T * x, const T * g, T * dst, const uint64_t k, const sycl::uint3 n_fd, const uint64_t o0, const uint64_t o1, const sycl::nd_item<1> &item_ct1) {
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
const sycl::uint2 rc = fast_div_modulo((uint32_t) i, n_fd);
const int64_t j0 = rc.x() * o0 + rc.y();
const int64_t j1 = o0 == o1 ? j0 : rc.x() * o1 + rc.y();
dst[i] = op_gelu_quick(x[j0]) * g[j1];
dst[i] = func(x[j0]) * g[j1];
}
}
@@ -670,6 +648,35 @@ static inline void ggml_sycl_op_unary(
});
}
template<typename F>
static inline void ggml_sycl_op_unary_gated(
ggml_backend_sycl_context & ctx, ggml_tensor * dst, F func) {
dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[func](const auto * x_ptr, const auto * g_ptr, auto * dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = (uint32_t) ceil_div(k, SYCL_GLU_BLOCK_SIZE);
const sycl::nd_range<1> launch_range(num_blocks * sycl::range<1>(SYCL_GLU_BLOCK_SIZE),
sycl::range<1>(SYCL_GLU_BLOCK_SIZE));
// o0 == n and o1 == n make the index math the identity, so index flat
// note: not ggml_is_contiguous - a fused [gate|up] src0 is contiguous with o0 == 2n
if (o0 == n && o1 == n) {
main_stream->parallel_for(launch_range,
[=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
unary_gated_op_flat_kernel(x_ptr, g_ptr, dst_ptr, k, item_ct1, func);
});
} else {
// launch-invariant divisor, and only this path needs it
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(launch_range,
[=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
unary_gated_op_generic_kernel(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1, func);
});
}
});
}
static inline void ggml_sycl_op_arange(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
GGML_ASSERT(dst->type == GGML_TYPE_F32);
@@ -967,42 +974,21 @@ static inline void ggml_sycl_op_acc(ggml_backend_sycl_context & ctx, ggml_tensor
}
static inline void ggml_sycl_op_geglu(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
ggml_sycl_detail::dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[](const auto* x_ptr, const auto* g_ptr, auto* dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = ceil_div(k, SYCL_GELU_BLOCK_SIZE);
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(
sycl::nd_range<1>((num_blocks * sycl::range<1>(SYCL_GELU_BLOCK_SIZE)),
sycl::range<1>(SYCL_GELU_BLOCK_SIZE)), [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
gated_op_fused_geglu(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1);
});
});
ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) {
return op_gelu(x);
});
}
static inline void ggml_sycl_op_reglu(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
ggml_sycl_detail::dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[](const auto* x_ptr, const auto* g_ptr, auto* dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = ceil_div((uint32_t)k, SYCL_RELU_BLOCK_SIZE); // Using RELU block size for reglu
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(
sycl::nd_range<1>((num_blocks * sycl::range<1>(SYCL_RELU_BLOCK_SIZE)),
sycl::range<1>(SYCL_RELU_BLOCK_SIZE)), [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
gated_op_fused_reglu(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1);
});
});
ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) {
return op_relu(x);
});
}
static inline void ggml_sycl_op_swiglu(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
ggml_sycl_detail::dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[](const auto* x_ptr, const auto* g_ptr, auto* dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = ceil_div((uint32_t)k, SYCL_SILU_BLOCK_SIZE); // Using SILU block size for swiglu
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(
sycl::nd_range<1>((num_blocks * sycl::range<1>(SYCL_SILU_BLOCK_SIZE)),
sycl::range<1>(SYCL_SILU_BLOCK_SIZE)), [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
gated_op_fused_swiglu(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1);
});
});
ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) {
return op_silu(x);
});
}
__dpct_inline__ float ggml_sycl_op_swiglu_oai_single(float x, float g, float alpha = 1.702f, float limit = 7.0f) {
@@ -1097,29 +1083,15 @@ void ggml_sycl_op_swiglu_oai(ggml_backend_sycl_context & ctx, ggml_tensor * dst)
}
static inline void ggml_sycl_op_geglu_erf(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
ggml_sycl_detail::dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[](const auto* x_ptr, const auto* g_ptr, auto* dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = ceil_div(k, SYCL_GELU_BLOCK_SIZE);
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(
sycl::nd_range<1>((num_blocks * sycl::range<1>(SYCL_GELU_BLOCK_SIZE)),
sycl::range<1>(SYCL_GELU_BLOCK_SIZE)), [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
gated_op_fused_geglu_erf(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1);
});
});
ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) {
return op_gelu_erf(x);
});
}
static inline void ggml_sycl_op_geglu_quick(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
ggml_sycl_detail::dispatch_ggml_sycl_op_fused_glu(ctx, dst,
[](const auto* x_ptr, const auto* g_ptr, auto* dst_ptr, uint64_t k, uint64_t n, uint64_t o0, uint64_t o1, queue_ptr main_stream) {
const uint32_t num_blocks = ceil_div(k, SYCL_GELU_BLOCK_SIZE);
const sycl::uint3 n_fd = init_fastdiv_values((uint32_t) n);
main_stream->parallel_for(
sycl::nd_range<1>((num_blocks * sycl::range<1>(SYCL_GELU_BLOCK_SIZE)),
sycl::range<1>(SYCL_GELU_BLOCK_SIZE)), [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
gated_op_fused_geglu_quick(x_ptr, g_ptr, dst_ptr, k, n_fd, o0, o1, item_ct1);
});
});
ggml_sycl_detail::ggml_sycl_op_unary_gated(ctx, dst, [](auto x) {
return op_gelu_quick(x);
});
}
+18 -16
View File
@@ -73,6 +73,7 @@ static void flash_attn_ext_vec(const char* __restrict__ Q,
const int32_t nb31,
const int32_t nb32,
const int64_t nb33) {
#ifdef SYCL_FLASH_ATTN
// Skip unused kernel variants for faster compilation:
@@ -469,7 +470,6 @@ static void flash_attn_ext_vec(const char* __restrict__ Q,
}
}
item_ct1.barrier(sycl::access::fence_space::local_space);
#pragma unroll
@@ -591,22 +591,24 @@ void ggml_sycl_flash_attn_ext_vec_case_impl(ggml_backend_sycl_context & ctx, ggm
const auto arch = ggml_sycl_info().devices[ctx.device].hw_info.arch;
const int nthreads = ggml_sycl_fattn_vec_get_nthreads_device(arch);
// 256 threads would overflow the 64 KB work-group local memory at D == 512, so keep 128 there.
if (D <= 256 && nthreads == 256) {
constexpr int nthreads_hw = 256;
constexpr int nwarps = nthreads_hw / warp_size;
launch_fattn<D, cols_per_block, 1,
flash_attn_ext_vec<D, cols_per_block, type_K, type_V,
use_logit_softcap, warp_size, nthreads_hw>, warp_size>(
ctx, dst, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false);
} else {
constexpr int nthreads_hw = 128;
constexpr int nwarps = nthreads_hw / warp_size;
launch_fattn<D, cols_per_block, 1,
flash_attn_ext_vec<D, cols_per_block, type_K, type_V,
use_logit_softcap, warp_size, nthreads_hw>, warp_size>(
ctx, dst, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false);
if constexpr (D <= 256) {
if (nthreads == 256) {
constexpr int nthreads_hw = 256;
constexpr int nwarps = nthreads_hw / warp_size;
launch_fattn<D, cols_per_block, 1,
flash_attn_ext_vec<D, cols_per_block, type_K, type_V,
use_logit_softcap, warp_size, nthreads_hw>, warp_size>(
ctx, dst, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false);
return;
}
}
constexpr int nthreads_hw = 128;
constexpr int nwarps = nthreads_hw / warp_size;
launch_fattn<D, cols_per_block, 1,
flash_attn_ext_vec<D, cols_per_block, type_K, type_V,
use_logit_softcap, warp_size, nthreads_hw>, warp_size>(
ctx, dst, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false);
}
template <int D, int type_K, int type_V>
+38 -8
View File
@@ -62,6 +62,8 @@
#include "ggml-sycl/repeat_back.hpp"
#include "ggml-sycl/set_rows.hpp"
#include "ggml-sycl/set.hpp"
#include "ggml-sycl/dsv4-hc.hpp"
#include "ggml-sycl/lightning-indexer.hpp"
#include "ggml-sycl/conv2d.hpp"
#include "ggml-sycl/conv2d-dw.hpp"
#include "ggml-sycl/conv2d-transpose.hpp"
@@ -4942,6 +4944,18 @@ static bool ggml_sycl_compute_forward(ggml_backend_sycl_context & ctx, struct gg
case GGML_OP_SET_ROWS:
ggml_sycl_op_set_rows(ctx, dst);
break;
case GGML_OP_DSV4_HC_PRE:
ggml_sycl_op_dsv4_hc_pre(ctx, dst);
break;
case GGML_OP_DSV4_HC_COMB:
ggml_sycl_op_dsv4_hc_comb(ctx, dst);
break;
case GGML_OP_DSV4_HC_POST:
ggml_sycl_op_dsv4_hc_post(ctx, dst);
break;
case GGML_OP_LIGHTNING_INDEXER:
ggml_sycl_op_lightning_indexer(ctx, dst);
break;
case GGML_OP_DUP:
ggml_sycl_dup(ctx, dst);
break;
@@ -5795,17 +5809,33 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons
case GGML_OP_SET_ROWS:
{
auto res = ((op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_BF16 ||
op->type == GGML_TYPE_Q8_0 || op->type == GGML_TYPE_Q5_1 || op->type == GGML_TYPE_Q5_0 ||
op->type == GGML_TYPE_Q1_0 ||
op->type == GGML_TYPE_Q4_1 || op->type == GGML_TYPE_Q4_0 || op->type == GGML_TYPE_IQ4_NL ||
op->type == GGML_TYPE_MXFP4 || op->type == GGML_TYPE_NVFP4) &&
op->src[0]->type == GGML_TYPE_F32 &&
(op->src[1]->type == GGML_TYPE_I64 || op->src[1]->type == GGML_TYPE_I32));
auto res = (op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_F16 ||
op->src[0]->type == GGML_TYPE_BF16) &&
(op->src[1]->type == GGML_TYPE_I64 || op->src[1]->type == GGML_TYPE_I32);
return res;
}
break;
case GGML_OP_DSV4_HC_PRE:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->type == GGML_TYPE_F32;
case GGML_OP_DSV4_HC_COMB:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
case GGML_OP_DSV4_HC_POST:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 &&
op->type == GGML_TYPE_F32;
case GGML_OP_LIGHTNING_INDEXER:
return op->src[0]->type == GGML_TYPE_F32 &&
(op->src[1]->type == GGML_TYPE_F16 || op->src[1]->type == GGML_TYPE_F32 ||
op->src[1]->type == GGML_TYPE_BF16 || op->src[1]->type == GGML_TYPE_Q8_0 ||
op->src[1]->type == GGML_TYPE_Q5_1 || op->src[1]->type == GGML_TYPE_Q5_0 ||
op->src[1]->type == GGML_TYPE_Q4_1 || op->src[1]->type == GGML_TYPE_Q4_0 ||
op->src[1]->type == GGML_TYPE_IQ4_NL) &&
op->src[2]->type == GGML_TYPE_F32 &&
op->src[3]->type == GGML_TYPE_F16 &&
op->type == GGML_TYPE_F32 &&
op->src[0]->ne[0] == WARP_SIZE * 8;
case GGML_OP_CPY:
{
ggml_type src0_type = op->src[0]->type;
+197
View File
@@ -0,0 +1,197 @@
#include "lightning-indexer.hpp"
#include "dequantize.hpp"
static void lightning_indexer_f32_sycl(
const char * q, const char * k, const char * w, const char * m, float * dst,
int64_t n_embd, int64_t n_head, int64_t n_batch, int64_t n_stream, int64_t n_kv,
int64_t nem3,
int64_t nbq1, int64_t nbq2, int64_t nbq3,
int64_t nbk2, int64_t nbk3,
int64_t nbw1, int64_t nbw3,
int64_t nbm1, int64_t nbm3,
int64_t nb1, int64_t nb3,
ggml_type k_type,
queue_ptr stream) {
constexpr int64_t LANES = WARP_SIZE;
constexpr int64_t ELEMS_PER_LANE = 8;
constexpr int64_t ROWS_PER_BLOCK = 4;
constexpr int64_t BLOCK_SIZE = ROWS_PER_BLOCK * LANES;
const int64_t n_rows = n_batch * n_stream * n_kv;
const int64_t n_blocks = (n_rows + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
stream->parallel_for(
sycl::nd_range<1>(
sycl::range<1>(n_blocks * BLOCK_SIZE),
sycl::range<1>(BLOCK_SIZE)),
[=](sycl::nd_item<1> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
const int64_t ir = item.get_global_id(0);
const int64_t lane = ir % LANES;
const int64_t row = ir / LANES;
if (row >= n_rows) {
return;
}
const int64_t i_bs = row / n_kv;
const int64_t i_kv = row % n_kv;
const int64_t i_batch = i_bs / n_stream;
const int64_t i_stream = i_bs % n_stream;
// load K row slice into registers (row is contiguous, nbk0 == type size)
const char * k_base = k + i_kv*nbk2 + i_stream*nbk3;
float k_local[ELEMS_PER_LANE];
if (k_type == GGML_TYPE_F16) {
const sycl::half * k_row = (const sycl::half *) k_base;
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
k_local[j] = static_cast<float>(k_row[lane*ELEMS_PER_LANE + j]);
}
} else if (k_type == GGML_TYPE_F32) {
const float * k_row = (const float *) k_base;
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
k_local[j] = k_row[lane*ELEMS_PER_LANE + j];
}
} else {
const int64_t lane_base = lane * ELEMS_PER_LANE;
switch (k_type) {
case GGML_TYPE_BF16: {
const sycl::ext::oneapi::bfloat16 * k_row = (const sycl::ext::oneapi::bfloat16 *) k_base;
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
k_local[j] = static_cast<float>(k_row[lane_base + j]);
}
} break;
case GGML_TYPE_Q4_0:
case GGML_TYPE_Q4_1:
case GGML_TYPE_Q5_0:
case GGML_TYPE_Q5_1: {
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
const int64_t idx = lane_base + j;
const int64_t ib = idx / QK4_0;
const int iqs = idx % (QK4_0/2);
dfloat2 kv;
if (k_type == GGML_TYPE_Q4_0) {
dequantize_q4_0(k_base, ib, iqs, kv);
} else if (k_type == GGML_TYPE_Q4_1) {
dequantize_q4_1(k_base, ib, iqs, kv);
} else if (k_type == GGML_TYPE_Q5_0) {
dequantize_q5_0(k_base, ib, iqs, kv);
} else {
dequantize_q5_1(k_base, ib, iqs, kv);
}
k_local[j] = (idx % QK4_0) < (QK4_0/2) ? static_cast<float>(kv.x()) : static_cast<float>(kv.y());
}
} break;
case GGML_TYPE_Q8_0: {
#pragma unroll
for (int64_t pair = 0; pair < ELEMS_PER_LANE / 2; ++pair) {
const int64_t elem0 = lane_base + 2 * pair;
dfloat2 kv;
dequantize_q8_0(k_base, elem0 / QK8_0, elem0 % QK8_0, kv);
k_local[2 * pair + 0] = static_cast<float>(kv.x());
k_local[2 * pair + 1] = static_cast<float>(kv.y());
}
} break;
case GGML_TYPE_IQ4_NL: {
#pragma unroll
for (int64_t pair = 0; pair < ELEMS_PER_LANE / 2; ++pair) {
const int64_t elem0 = lane_base + 2 * pair;
dfloat2 kv;
dequantize_iq4_nl(k_base, elem0 / QK4_NL, elem0 % QK4_NL, kv);
k_local[2 * pair + 0] = static_cast<float>(kv.x());
k_local[2 * pair + 1] = static_cast<float>(kv.y());
}
} break;
default:
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
k_local[j] = 0.0f;
}
break;
}
}
const char * q_base = q + i_batch*nbq2 + i_stream*nbq3;
const float * w_base = (const float *) (w + i_batch*nbw1 + i_stream*nbw3);
float score = 0.0f;
for (int64_t h = 0; h < n_head; ++h) {
const float * q_row = (const float *) (q_base + h*nbq1);
float dot = 0.0f;
#pragma unroll
for (int64_t j = 0; j < ELEMS_PER_LANE; ++j) {
const int64_t i = lane*ELEMS_PER_LANE + j;
if (i < n_embd) {
dot += q_row[i] * k_local[j];
}
}
dot = sycl::reduce_over_group(item.get_sub_group(), dot, sycl::plus<float>());
if (lane == 0) {
score += sycl::max(dot, 0.0f) * w_base[h];
}
}
if (lane == 0) {
const sycl::half * m_base = (const sycl::half *) (m + i_batch*nbm1 + (i_stream % nem3)*nbm3);
// flat-index store: storing through a strided base pointer
// hangs/misroutes writes on this stack when n_batch*n_stream > 1
const int64_t dst_idx = i_kv + i_batch*(nb1/sizeof(float)) + i_stream*(nb3/sizeof(float));
dst[dst_idx] = score + static_cast<float>(m_base[i_kv]);
}
});
}
void ggml_sycl_op_lightning_indexer(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
scope_op_debug_print scope_dbg_print(__func__, dst, /*num_src=*/4);
const ggml_tensor * q = dst->src[0];
const ggml_tensor * k = dst->src[1];
const ggml_tensor * w = dst->src[2]; // weights
const ggml_tensor * m = dst->src[3]; // mask
GGML_ASSERT(dst->type == GGML_TYPE_F32);
GGML_ASSERT( q->type == GGML_TYPE_F32);
GGML_ASSERT( w->type == GGML_TYPE_F32);
GGML_ASSERT( m->type == GGML_TYPE_F16);
GGML_ASSERT(k->type == GGML_TYPE_F16 || k->type == GGML_TYPE_F32 || k->type == GGML_TYPE_BF16 ||
k->type == GGML_TYPE_Q8_0 || k->type == GGML_TYPE_Q5_1 || k->type == GGML_TYPE_Q5_0 ||
k->type == GGML_TYPE_Q4_1 || k->type == GGML_TYPE_Q4_0 || k->type == GGML_TYPE_IQ4_NL);
GGML_TENSOR_LOCALS(int64_t, neq, q, ne);
GGML_TENSOR_LOCALS(size_t, nbq, q, nb);
GGML_TENSOR_LOCALS(int64_t, nek, k, ne);
GGML_TENSOR_LOCALS(size_t, nbk, k, nb);
GGML_TENSOR_LOCALS(size_t, nbw, w, nb);
GGML_TENSOR_LOCALS(int64_t, nem, m, ne);
GGML_TENSOR_LOCALS(size_t, nbm, m, nb);
GGML_TENSOR_LOCALS(int64_t, ne, dst, ne);
GGML_TENSOR_LOCALS(size_t, nb, dst, nb);
// input rows must be contiguous
GGML_ASSERT(nbq0 == ggml_type_size(q->type));
GGML_ASSERT(nbk0 == ggml_type_size(k->type));
GGML_ASSERT(nbm0 == ggml_type_size(m->type));
GGML_ASSERT(nb0 == ggml_type_size(dst->type));
const int64_t n_embd = neq0;
const int64_t n_head = neq1;
const int64_t n_batch = neq2;
const int64_t n_stream = neq3;
const int64_t n_kv = nek2;
GGML_ASSERT(n_embd == WARP_SIZE * 8);
lightning_indexer_f32_sycl(
(const char *) q->data, (const char *) k->data,
(const char *) w->data, (const char *) m->data, (float *) dst->data,
n_embd, n_head, n_batch, n_stream, n_kv, nem3,
nbq1, nbq2, nbq3,
nbk2, nbk3,
nbw1, nbw3,
nbm1, nbm3,
nb1, nb3,
k->type,
ctx.stream());
}
+8
View File
@@ -0,0 +1,8 @@
#ifndef GGML_SYCL_LIGHTNING_INDEXER_HPP
#define GGML_SYCL_LIGHTNING_INDEXER_HPP
#include "common.hpp"
void ggml_sycl_op_lightning_indexer(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
#endif // GGML_SYCL_LIGHTNING_INDEXER_HPP
-2
View File
@@ -20,8 +20,6 @@
#define MATRIX_ROW_PADDING 512 // last row of quant. matrices is a multiple of this to avoid out-of-bounds memory accesses
#define SYCL_COL2IM_1D_BLOCK_SIZE 256
#define SYCL_GELU_BLOCK_SIZE 256
#define SYCL_SILU_BLOCK_SIZE 256
#define SYCL_TANH_BLOCK_SIZE 256
#define SYCL_RELU_BLOCK_SIZE 256
#define SYCL_HARDSIGMOID_BLOCK_SIZE 256
+344 -16
View File
@@ -1,6 +1,10 @@
#include "set_rows.hpp"
#include "cpy.hpp"
#include "ggml-quants.h"
#include <vector>
namespace utils {
template<typename T>
static constexpr bool is_arithmetic_v() {
@@ -20,7 +24,17 @@ convert (const char* src, char* dst) {
*reinterpret_cast<TOut*>(dst) = dst_val;
}
template <typename TIdx, typename blockType, int qk, cpy_kernel_t cpyblck>
#ifdef GGML_SYCL_HAS_BF16
// sycl::vec::convert does not provide a half -> bfloat16 path, so route through float.
template<>
inline void convert<sycl::half, sycl::ext::oneapi::bfloat16>(const char* src, char* dst) {
const float tmp = sycl::vec<sycl::half, 1>(*reinterpret_cast<const sycl::half*>(src))
.template convert<float, sycl::rounding_mode::automatic>()[0];
*reinterpret_cast<sycl::ext::oneapi::bfloat16*>(dst) = sycl::ext::oneapi::bfloat16(tmp);
}
#endif
template <typename TIn, typename TIdx, typename blockType, int qk, cpy_kernel_t cpyblck>
static void set_rows_sycl_q(const char * __restrict__ src0_d,
const TIdx * __restrict__ src1_d,
blockType * __restrict__ dst_d,
@@ -68,13 +82,22 @@ static void set_rows_sycl_q(const char * __restrict__ src0_d,
const int64_t i11 = i02 % ne11;
const int64_t i10 = i01;
const size_t src_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 });
const char * src_block = src0_d + src_offset + i00 * sizeof(float);
const char * src_block = src0_d + src_offset + i00 * sizeof(TIn);
const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 });
const int64_t dst_row = src1_d[src1_offset / sizeof(TIdx)];
const size_t dst_offset =
calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 }) + (i00 / qk) * sizeof(blockType);
char * dst_block = reinterpret_cast<char *>(reinterpret_cast<char *>(dst_d) + dst_offset);
cpyblck(src_block, dst_block);
if constexpr (std::is_same_v<TIn, float>) {
cpyblck(src_block, dst_block);
} else {
float src_block_f32[qk];
const TIn * src_block_t = reinterpret_cast<const TIn *>(src_block);
for (int j = 0; j < qk; ++j) {
src_block_f32[j] = (float) src_block_t[j];
}
cpyblck(reinterpret_cast<const char *>(src_block_f32), dst_block);
}
});
GGML_UNUSED(ne10);
GGML_UNUSED(ne13);
@@ -82,6 +105,139 @@ static void set_rows_sycl_q(const char * __restrict__ src0_d,
GGML_UNUSED(nb13);
}
template<typename blockType>
using quantize_row_qk_t = void (*)(const float *, blockType *, int64_t);
using quantize_rows_f_t = size_t (*)(const float *, void *, int64_t, int64_t, const float *);
template <typename TIn, typename TIdx, typename blockType, int qk, quantize_row_qk_t<blockType> quantize_row>
static void set_rows_sycl_qk_host(
const ggml_tensor * src0,
const ggml_tensor * src1,
ggml_tensor * dst,
const int64_t ne00,
const int64_t ne01,
const int64_t ne02,
const int64_t ne03,
const int64_t ne11,
const int64_t ne12,
const size_t nb01,
const size_t nb02,
const size_t nb03,
const size_t nb10,
const size_t nb11,
const size_t nb12,
const size_t nb1,
const size_t nb2,
const size_t nb3,
queue_ptr stream) {
GGML_ASSERT(ne00 % qk == 0);
const size_t src0_bytes = ggml_nbytes(src0);
const size_t src1_bytes = ggml_nbytes(src1);
std::vector<char> src0_host(src0_bytes);
std::vector<char> src1_host(src1_bytes);
stream->memcpy(src0_host.data(), src0->data, src0_bytes);
stream->memcpy(src1_host.data(), src1->data, src1_bytes);
stream->wait();
std::vector<float> src_row_f32(ne00);
const int64_t nblocks = ne00 / qk;
std::vector<blockType> dst_row_q(nblocks);
for (int64_t i03 = 0; i03 < ne03; ++i03) {
for (int64_t i02 = 0; i02 < ne02; ++i02) {
for (int64_t i01 = 0; i01 < ne01; ++i01) {
const int64_t i12 = i03 % ne12;
const int64_t i11 = i02 % ne11;
const int64_t i10 = i01;
const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 });
const int64_t dst_row = *(const TIdx *) (src1_host.data() + src1_offset);
const size_t src0_row_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 });
const TIn * src_row = reinterpret_cast<const TIn *>(src0_host.data() + src0_row_offset);
for (int64_t i00 = 0; i00 < ne00; ++i00) {
src_row_f32[i00] = (float) src_row[i00];
}
quantize_row(src_row_f32.data(), dst_row_q.data(), ne00);
const size_t dst_offset = calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 });
stream->memcpy((char *) dst->data + dst_offset, dst_row_q.data(), nblocks * sizeof(blockType));
stream->wait();
}
}
}
}
template <typename TIn, typename TIdx, typename blockType, int qk, quantize_rows_f_t quantize_rows>
static void set_rows_sycl_iq_host(
const ggml_tensor * src0,
const ggml_tensor * src1,
ggml_tensor * dst,
const int64_t ne00,
const int64_t ne01,
const int64_t ne02,
const int64_t ne03,
const int64_t ne11,
const int64_t ne12,
const size_t nb01,
const size_t nb02,
const size_t nb03,
const size_t nb10,
const size_t nb11,
const size_t nb12,
const size_t nb1,
const size_t nb2,
const size_t nb3,
queue_ptr stream) {
GGML_ASSERT(ne00 % qk == 0);
const size_t src0_bytes = ggml_nbytes(src0);
const size_t src1_bytes = ggml_nbytes(src1);
std::vector<char> src0_host(src0_bytes);
std::vector<char> src1_host(src1_bytes);
stream->memcpy(src0_host.data(), src0->data, src0_bytes);
stream->memcpy(src1_host.data(), src1->data, src1_bytes);
stream->wait();
std::vector<float> src_row_f32(ne00);
const int64_t nblocks = ne00 / qk;
std::vector<blockType> dst_row_q(nblocks);
for (int64_t i03 = 0; i03 < ne03; ++i03) {
for (int64_t i02 = 0; i02 < ne02; ++i02) {
for (int64_t i01 = 0; i01 < ne01; ++i01) {
const int64_t i12 = i03 % ne12;
const int64_t i11 = i02 % ne11;
const int64_t i10 = i01;
const size_t src1_offset = calculate_offset<3>({ nb10, nb11, nb12 }, { i10, i11, i12 });
const int64_t dst_row = *(const TIdx *) (src1_host.data() + src1_offset);
const size_t src0_row_offset = calculate_offset<3>({ nb01, nb02, nb03 }, { i01, i02, i03 });
const TIn * src_row = reinterpret_cast<const TIn *>(src0_host.data() + src0_row_offset);
for (int64_t i00 = 0; i00 < ne00; ++i00) {
src_row_f32[i00] = (float) src_row[i00];
}
quantize_rows(src_row_f32.data(), dst_row_q.data(), 1, ne00, nullptr);
const size_t dst_offset = calculate_offset<3>({ nb1, nb2, nb3 }, { dst_row, i02, i03 });
stream->memcpy((char *) dst->data + dst_offset, dst_row_q.data(), nblocks * sizeof(blockType));
stream->wait();
}
}
}
}
template<typename TIn, typename TIdx, typename TOut>
static void k_set_rows(
const char * __restrict__ src0, const TIdx * __restrict__ src1, char * __restrict__ dst,
@@ -200,31 +356,194 @@ static void set_rows_sycl(ggml_backend_sycl_context & ctx, const ggml_tensor * s
break;
#endif
case GGML_TYPE_Q8_0:
set_rows_sycl_q<TIdx, block_q8_0, QK8_0, cpy_blck_f32_q8_0>(src0_d, src1_d, (block_q8_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q8_0, QK8_0, cpy_blck_f32_q8_0>(
src0_d, src1_d, (block_q8_0 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q1_0:
set_rows_sycl_q<TIdx, block_q1_0, QK1_0, cpy_blck_f32_q1_0>(src0_d, src1_d, (block_q1_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q1_0, QK1_0, cpy_blck_f32_q1_0>(
src0_d, src1_d, (block_q1_0 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q2_0:
set_rows_sycl_q<TIn, TIdx, block_q2_0, QK2_0, cpy_blck_f32_q2_0>(
src0_d, src1_d, (block_q2_0 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q5_1:
set_rows_sycl_q<TIdx, block_q5_1, QK5_1, cpy_blck_f32_q5_1>(src0_d, src1_d, (block_q5_1 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q5_1, QK5_1, cpy_blck_f32_q5_1>(
src0_d, src1_d, (block_q5_1 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q5_0:
set_rows_sycl_q<TIdx, block_q5_0, QK5_0, cpy_blck_f32_q5_0>(src0_d, src1_d, (block_q5_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q5_0, QK5_0, cpy_blck_f32_q5_0>(
src0_d, src1_d, (block_q5_0 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q4_1:
set_rows_sycl_q<TIdx, block_q4_1, QK4_1, cpy_blck_f32_q4_1>(src0_d, src1_d, (block_q4_1 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q4_1, QK4_1, cpy_blck_f32_q4_1>(
src0_d, src1_d, (block_q4_1 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q4_0:
set_rows_sycl_q<TIdx, block_q4_0, QK4_0, cpy_blck_f32_q4_0>(src0_d, src1_d, (block_q4_0 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_q4_0, QK4_0, cpy_blck_f32_q4_0>(
src0_d, src1_d, (block_q4_0 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_IQ4_NL:
set_rows_sycl_q<TIdx, block_iq4_nl, QK4_NL, cpy_blck_f32_iq4_nl>(src0_d, src1_d, (block_iq4_nl *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_iq4_nl, QK4_NL, cpy_blck_f32_iq4_nl>(
src0_d, src1_d, (block_iq4_nl *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_MXFP4:
set_rows_sycl_q<TIdx, block_mxfp4, QK_MXFP4, cpy_blck_f32_mxfp4>(src0_d, src1_d, (block_mxfp4 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_mxfp4, QK_MXFP4, cpy_blck_f32_mxfp4>(
src0_d, src1_d, (block_mxfp4 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_NVFP4:
set_rows_sycl_q<TIdx, block_nvfp4, QK_NVFP4, cpy_blck_f32_nvfp4>(src0_d, src1_d, (block_nvfp4 *)dst->data, ne00, ne01, ne02, ne03, ne10, ne11, ne12, ne13, nb00, nb01, nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
set_rows_sycl_q<TIn, TIdx, block_nvfp4, QK_NVFP4, cpy_blck_f32_nvfp4>(
src0_d, src1_d, (block_nvfp4 *) dst->data, ne00, ne01, ne02, ne03,
ne10, ne11, ne12, ne13, nb00, nb01,
nb02, nb03, nb10, nb11, nb12, nb13, nb1, nb2, nb3, stream);
break;
case GGML_TYPE_Q2_K:
set_rows_sycl_qk_host<TIn, TIdx, block_q2_K, QK_K, quantize_row_q2_K_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_Q3_K:
set_rows_sycl_qk_host<TIn, TIdx, block_q3_K, QK_K, quantize_row_q3_K_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_Q4_K:
set_rows_sycl_qk_host<TIn, TIdx, block_q4_K, QK_K, quantize_row_q4_K_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_Q5_K:
set_rows_sycl_qk_host<TIn, TIdx, block_q5_K, QK_K, quantize_row_q5_K_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_Q6_K:
set_rows_sycl_qk_host<TIn, TIdx, block_q6_K, QK_K, quantize_row_q6_K_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ2_XXS:
set_rows_sycl_iq_host<TIn, TIdx, block_iq2_xxs, QK_K, quantize_iq2_xxs>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ2_XS:
set_rows_sycl_iq_host<TIn, TIdx, block_iq2_xs, QK_K, quantize_iq2_xs>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ2_S:
set_rows_sycl_iq_host<TIn, TIdx, block_iq2_s, QK_K, quantize_iq2_s>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ3_XXS:
set_rows_sycl_qk_host<TIn, TIdx, block_iq3_xxs, QK_K, quantize_row_iq3_xxs_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ3_S:
set_rows_sycl_qk_host<TIn, TIdx, block_iq3_s, QK_K, quantize_row_iq3_s_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ1_S:
set_rows_sycl_iq_host<TIn, TIdx, block_iq1_s, QK_K, quantize_iq1_s>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ1_M:
set_rows_sycl_iq_host<TIn, TIdx, block_iq1_m, QK_K, quantize_iq1_m>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
case GGML_TYPE_IQ4_XS:
set_rows_sycl_qk_host<TIn, TIdx, block_iq4_xs, QK_K, quantize_row_iq4_xs_ref>(
src0, src1, dst,
ne00, ne01, ne02, ne03,
ne11, ne12,
nb01, nb02, nb03,
nb10, nb11, nb12,
nb1, nb2, nb3,
stream);
break;
default:
GGML_ABORT("Unsupported tensor type!");
@@ -237,12 +556,21 @@ void ggml_sycl_op_set_rows(ggml_backend_sycl_context & ctx, ggml_tensor * dst) {
const ggml_tensor * src0 = dst->src[0];
const ggml_tensor * src1 = dst->src[1];
GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32);
GGML_ASSERT(dst->src[0]->type == GGML_TYPE_F32 || dst->src[0]->type == GGML_TYPE_F16);
GGML_ASSERT(dst->src[1]->type == GGML_TYPE_I64 || dst->src[1]->type == GGML_TYPE_I32);
if (src1->type == GGML_TYPE_I64) {
set_rows_sycl<float, int64_t>(ctx, src0, src1, dst);
// dispatch on the index type (src1) and the source value type (src0)
if (src0->type == GGML_TYPE_F16) {
if (src1->type == GGML_TYPE_I64) {
set_rows_sycl<sycl::half, int64_t>(ctx, src0, src1, dst);
} else {
set_rows_sycl<sycl::half, int32_t>(ctx, src0, src1, dst);
}
} else {
set_rows_sycl<float, int32_t>(ctx, src0, src1, dst);
if (src1->type == GGML_TYPE_I64) {
set_rows_sycl<float, int64_t>(ctx, src0, src1, dst);
} else {
set_rows_sycl<float, int32_t>(ctx, src0, src1, dst);
}
}
}
+7 -3
View File
@@ -36,9 +36,13 @@ static void kernel_ssm_conv(
return;
}
const int channel = static_cast<int>(idx % d_inner);
const int token = static_cast<int>((idx / d_inner) % n_t);
const int seq = static_cast<int>(idx / (static_cast<size_t>(d_inner) * static_cast<size_t>(n_t)));
// src has the tokens of one channel contiguous, dst has the channels of one
// token contiguous, so either the loads or the store must be strided. Indexing
// token-fastest coalesces the d_conv loads, which measured faster except for
// short, cache-resident rows.
const int token = static_cast<int>(idx % n_t);
const int channel = static_cast<int>((idx / n_t) % d_inner);
const int seq = static_cast<int>(idx / (static_cast<size_t>(n_t) * static_cast<size_t>(d_inner)));
const float *s = src_data
+ static_cast<size_t>(seq) * static_cast<size_t>(src_stride_seq)
+1 -1
View File
@@ -1 +1 @@
90951f99af1fbebef3fbdd58ff5b8715b0bb9c43
30bf8685ed4eb0a47f2b06229543327749904150
+16
View File
@@ -8722,6 +8722,13 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_l2_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, true));
test_cases.emplace_back(new test_l2_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, eps, false, true));
}
// row lengths that are not a multiple of 32, for the scalar (33) and float4 (132, 260) paths
for (uint32_t n : { 33, 132, 260 }) {
for (bool v : { false, true }) {
test_cases.emplace_back(new test_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, v, eps));
test_cases.emplace_back(new test_rms_norm(GGML_TYPE_F32, { n, 5, 4, 3 }, v, eps));
}
}
}
// in-place tests
@@ -9747,6 +9754,15 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
std::vector<std::unique_ptr<test_case>> test_cases;
// SWIGLU at a 27B-class FFN width, fused [gate|up] vs split operands
// note: same bytes either way, so a backend that indexes them differently shows it here
for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) {
for (int64_t n_tokens : {512, 2048}) {
test_cases.emplace_back(new test_glu(GGML_GLU_OP_SWIGLU, type, { 2*17408, n_tokens, 1, 1 }, 0, false));
test_cases.emplace_back(new test_glu_split(GGML_GLU_OP_SWIGLU, type, { 17408, n_tokens, 1, 1 }, 0));
}
}
// Conv2d: K=CRS=NPQ=4096 matmul performance
uint32_t iwh_idx = 0;
uint32_t kwh_idx = 1;
+1 -1
View File
@@ -195,7 +195,7 @@ static const std::vector<std::string> dspark_dflash = {
struct plan_case {
const char * name;
const std::vector<std::string> & files;
const std::vector<std::string> files;
const char * hf_repo;
const char * hf_file;
bool sidecars; // request mmproj + mtp + dflash + eagle3 + dspark
+68
View File
@@ -1,4 +1,6 @@
#include <stdio.h>
#include <stdlib.h>
#include <string.h>
#include <assert.h>
#include "mtmd.h"
@@ -62,6 +64,72 @@ int main(void) {
}
}
// test chunk save/load round-trip
for (size_t i = 0; i < n_chunks; i++) {
const mtmd_input_chunk * chunk = mtmd_input_chunks_get(chunks, i);
assert(chunk != NULL);
enum mtmd_input_chunk_type type = mtmd_input_chunk_get_type(chunk);
// query the required buffer size (out_buf == NULL)
size_t expected_len = 0;
int32_t rc = mtmd_input_chunk_save(chunk, NULL, 0, &expected_len);
printf(" Chunk %zu: save query rc = %d, expected_len = %zu\n", i, rc, expected_len);
assert(rc == 0);
assert(expected_len > 0);
// saving into a too-small buffer must fail, not crash
char tiny_buf[1];
rc = mtmd_input_chunk_save(chunk, tiny_buf, sizeof(tiny_buf), NULL);
printf(" Chunk %zu: save into too-small buffer rc = %d (expect non-zero)\n", i, rc);
assert(rc != 0);
// save into a properly-sized buffer
char * buf = (char *) malloc(expected_len);
assert(buf != NULL);
rc = mtmd_input_chunk_save(chunk, buf, expected_len, NULL);
assert(rc == 0);
// loading from a truncated buffer must fail gracefully, not crash
if (expected_len > 1) {
mtmd_input_chunk * bad = mtmd_input_chunk_load(buf, expected_len - 1);
printf(" Chunk %zu: load from truncated buffer = %p (expect NULL)\n", i, (void *) bad);
assert(bad == NULL);
}
// load it back
mtmd_input_chunk * loaded = mtmd_input_chunk_load(buf, expected_len);
assert(loaded != NULL);
// metadata must match the original chunk
assert(mtmd_input_chunk_get_type(loaded) == type);
assert(mtmd_input_chunk_get_n_tokens(loaded) == mtmd_input_chunk_get_n_tokens(chunk));
assert(mtmd_input_chunk_get_n_pos(loaded) == mtmd_input_chunk_get_n_pos(chunk));
if (type == MTMD_INPUT_CHUNK_TYPE_TEXT) {
size_t n_tok_orig, n_tok_loaded;
const llama_token * tok_orig = mtmd_input_chunk_get_tokens_text(chunk, &n_tok_orig);
const llama_token * tok_loaded = mtmd_input_chunk_get_tokens_text(loaded, &n_tok_loaded);
printf(" Chunk %zu: loaded %zu text tokens (orig %zu), first token %d (orig %d)\n",
i, n_tok_loaded, n_tok_orig,
n_tok_loaded > 0 ? tok_loaded[0] : -1,
n_tok_orig > 0 ? tok_orig[0] : -1);
assert(n_tok_orig == n_tok_loaded);
for (size_t j = 0; j < n_tok_orig; j++) {
assert(tok_orig[j] == tok_loaded[j]);
}
} else if (type == MTMD_INPUT_CHUNK_TYPE_IMAGE || type == MTMD_INPUT_CHUNK_TYPE_AUDIO) {
const char * id_orig = mtmd_input_chunk_get_id(chunk);
const char * id_loaded = mtmd_input_chunk_get_id(loaded);
printf(" Chunk %zu: loaded id '%s' (orig '%s')\n", i, id_loaded, id_orig);
assert(id_orig != NULL && id_loaded != NULL);
assert(strcmp(id_orig, id_loaded) == 0);
}
mtmd_input_chunk_free(loaded);
free(buf);
}
printf("Chunk save/load round-trip OK\n");
// Free the chunks
mtmd_input_chunks_free(chunks);
+8
View File
@@ -591,6 +591,8 @@ struct clip_image_u8 {
}
};
struct mtmd_serialization; // forward declaration
// For images, buf.size() == nx*ny*3
// Memory layout: RGBRGBRGB...
// For seq, buf.size() == nx*ny*3*nt
@@ -671,6 +673,9 @@ struct clip_image_f32 {
return buf.empty();
}
void serialize(struct mtmd_serialization & ser) const;
void deserialize(struct mtmd_serialization & ser);
private:
std::vector<float> buf;
int nx_ = 0;
@@ -752,6 +757,9 @@ struct clip_image_f32_batch {
}
return new_batch;
}
void serialize(struct mtmd_serialization & ser) const;
void deserialize(struct mtmd_serialization & ser);
};
//
+11
View File
@@ -170,6 +170,17 @@ struct clip_hparams {
warmup_image_size = static_cast<int>(std::sqrt(image_max_pixels));
}
// used by longest_edge preprocessor (no model-specific value for min/max tokens)
void set_limit_image_tokens() {
const int patch_area = patch_size * patch_size * n_merge * n_merge;
if (custom_image_min_tokens > 0) {
image_min_pixels = custom_image_min_tokens * patch_area;
}
if (custom_image_max_tokens > 0) {
image_max_pixels = custom_image_max_tokens * patch_area;
}
}
void set_warmup_n_tokens(int n_tokens) {
int n_tok_per_side = static_cast<int>(std::sqrt(n_tokens));
GGML_ASSERT(n_tok_per_side * n_tok_per_side == n_tokens && "n_tokens must be n*n");
+3
View File
@@ -1434,6 +1434,7 @@ struct clip_model_loader {
// use default llava-uhd preprocessing params
get_u32(KEY_PROJ_SCALE_FACTOR, hparams.n_merge, false);
get_u32(KEY_PREPROC_IMAGE_SIZE, hparams.image_longest_edge, false);
hparams.set_limit_image_tokens();
} break;
case PROJECTOR_TYPE_LFM2:
{
@@ -1471,6 +1472,7 @@ struct clip_model_loader {
get_u32(KEY_SPATIAL_MERGE_SIZE, hparams.n_merge, false);
hparams.image_longest_edge = hparams.image_size;
get_u32(KEY_PREPROC_IMAGE_SIZE, hparams.image_longest_edge, false);
hparams.set_limit_image_tokens();
hparams.set_warmup_n_tokens(256); // avoid OOM on warmup
} break;
case PROJECTOR_TYPE_DOTS_OCR:
@@ -1595,6 +1597,7 @@ struct clip_model_loader {
if (hparams.image_longest_edge == 0) {
hparams.image_longest_edge = 3024;
}
// note: the step3vl preprocessor slices based on a fixed window grid, so it does not support custom min/max image tokens
hparams.warmup_image_size = hparams.image_size;
} break;
case PROJECTOR_TYPE_YOUTUVL:
+5 -11
View File
@@ -112,7 +112,6 @@ public:
c2w_state.clear();
audio_pcm.clear();
overlay.clear();
overlay_idx = 0;
h_state_buf.clear();
out_buf.clear();
prompt_embd_buf.clear();
@@ -205,11 +204,9 @@ public:
top_p = inp->top_p > 0 ? inp->top_p : 1.0f;
out_type = inp->out_type;
// the text stream keeps flowing during generation: after frame k, the input adds
// trailing text row k on top of the codes embedding, then tts_eos, then tts_pad
for (int i = 3; i < n_ids - 5; i++) overlay.push_back(row(ids[(size_t) i]));
overlay.push_back(row(tts_eos));
overlay.push_back(row(tts_pad));
// the prompt above holds the whole text stream up to tts_eos, so every generated
// frame adds tts_pad on top of the codes embedding
overlay = row(tts_pad);
return 0;
}
@@ -265,9 +262,7 @@ public:
}
std::vector<float> fb(out.embd, out.embd + n_embd);
const auto & ov = overlay[std::min(overlay_idx, overlay.size() - 1)];
for (int i = 0; i < n_embd; i++) fb[(size_t) i] += ov[(size_t) i];
overlay_idx++;
for (int i = 0; i < n_embd; i++) fb[(size_t) i] += overlay[(size_t) i];
const int n_pos_per_embd = mrope ? 4 : 1;
decode_embd_batch batch_embd(fb.data(), 1, n_pos_per_embd, n_embd);
@@ -437,8 +432,7 @@ private:
std::vector<int32_t> codes_buf;
std::vector<uint8_t> c2w_state;
std::vector<float> audio_pcm;
std::vector<std::vector<float>> overlay;
size_t overlay_idx = 0;
std::vector<float> overlay;
std::vector<float> h_state_buf;
mtmd_helper_gen_audio_outtype out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
std::vector<char> out_buf;
+51 -47
View File
@@ -139,50 +139,46 @@ struct img_tool {
}
}
// calculate the size of the **resized** image, while preserving the aspect ratio
// the calculated size will be aligned to the nearest multiple of align_size
// if H or W size is larger than longest_edge, it will be resized to longest_edge
static clip_image_size calc_size_preserved_ratio(const clip_image_size & inp_size, const int align_size, const int longest_edge) {
GGML_ASSERT(align_size > 0);
if (inp_size.width <= 0 || inp_size.height <= 0 || longest_edge <= 0) {
struct calc_size_opt {
int align_size = 1;
int min_pixels = 0; // 0 = disabled
int max_pixels = 0; // 0 = disabled
// applied before min/max_pixels, so min_pixels can push an edge back above longest_edge
int longest_edge = 0; // 0 = disabled
};
// calculate the size of the **resized** image, while preserving the aspect ratio and
// aligning to the nearest multiple of align_size ("smart_resize" in transformers code)
static clip_image_size calc_size_preserved_ratio(const clip_image_size & inp_size, const calc_size_opt & opts) {
GGML_ASSERT(opts.align_size > 0);
const int width = inp_size.width;
const int height = inp_size.height;
if (width <= 0 || height <= 0) {
return {0, 0};
}
float scale = std::min(static_cast<float>(longest_edge) / inp_size.width,
static_cast<float>(longest_edge) / inp_size.height);
auto round_by_factor = [f = opts.align_size](float x) { return static_cast<int>(std::round(x / static_cast<float>(f))) * f; };
auto ceil_by_factor = [f = opts.align_size](float x) { return static_cast<int>(std::ceil(x / static_cast<float>(f))) * f; };
auto floor_by_factor = [f = opts.align_size](float x) { return static_cast<int>(std::floor(x / static_cast<float>(f))) * f; };
float target_width_f = static_cast<float>(inp_size.width) * scale;
float target_height_f = static_cast<float>(inp_size.height) * scale;
int w_bar, h_bar;
if (opts.longest_edge > 0) {
const float scale = std::min(static_cast<float>(opts.longest_edge) / width,
static_cast<float>(opts.longest_edge) / height);
w_bar = ceil_by_factor(width * scale);
h_bar = ceil_by_factor(height * scale);
} else {
// always align up first
w_bar = std::max(opts.align_size, round_by_factor(width));
h_bar = std::max(opts.align_size, round_by_factor(height));
}
auto ceil_by_factor = [f = align_size](float x) { return static_cast<int>(std::ceil(x / static_cast<float>(f))) * f; };
int aligned_width = ceil_by_factor(target_width_f);
int aligned_height = ceil_by_factor(target_height_f);
return {aligned_width, aligned_height};
}
// calculate the size of the **resized** image, while preserving the aspect ratio
// the calculated size will have min_pixels <= W*H <= max_pixels
// this is referred as "smart_resize" in transformers code
static clip_image_size calc_size_preserved_ratio(const clip_image_size & inp_size, const int align_size, const int min_pixels, const int max_pixels) {
GGML_ASSERT(align_size > 0);
const int width = inp_size.width;
const int height = inp_size.height;
auto round_by_factor = [f = align_size](float x) { return static_cast<int>(std::round(x / static_cast<float>(f))) * f; };
auto ceil_by_factor = [f = align_size](float x) { return static_cast<int>(std::ceil(x / static_cast<float>(f))) * f; };
auto floor_by_factor = [f = align_size](float x) { return static_cast<int>(std::floor(x / static_cast<float>(f))) * f; };
// always align up first
int h_bar = std::max(align_size, round_by_factor(height));
int w_bar = std::max(align_size, round_by_factor(width));
if (h_bar * w_bar > max_pixels) {
const auto beta = std::sqrt(static_cast<float>(height * width) / max_pixels);
h_bar = std::max(align_size, floor_by_factor(height / beta));
w_bar = std::max(align_size, floor_by_factor(width / beta));
} else if (h_bar * w_bar < min_pixels) {
const auto beta = std::sqrt(static_cast<float>(min_pixels) / (height * width));
if (opts.max_pixels > 0 && h_bar * w_bar > opts.max_pixels) {
const auto beta = std::sqrt(static_cast<float>(height) * width / opts.max_pixels);
h_bar = std::max(opts.align_size, floor_by_factor(height / beta));
w_bar = std::max(opts.align_size, floor_by_factor(width / beta));
} else if (opts.min_pixels > 0 && h_bar * w_bar < opts.min_pixels) {
const auto beta = std::sqrt(static_cast<float>(opts.min_pixels) / (static_cast<float>(height) * width));
h_bar = ceil_by_factor(height * beta);
w_bar = ceil_by_factor(width * beta);
}
@@ -937,9 +933,12 @@ mtmd_image_preproc_out mtmd_image_preprocessor_dyn_size::preprocess(const clip_i
const int cur_merge = hparams.n_merge;
const clip_image_size target_size = img_tool::calc_size_preserved_ratio(
original_size,
hparams.patch_size * cur_merge,
hparams.image_min_pixels,
hparams.image_max_pixels);
{
/* align_size */ hparams.patch_size * cur_merge,
/* min_pixels */ hparams.image_min_pixels,
/* max_pixels */ hparams.image_max_pixels,
/* longest_edge */ 0,
});
img_tool::resize(img, resized_image, target_size,
hparams.image_resize_algo,
hparams.image_resize_pad,
@@ -961,8 +960,12 @@ mtmd_image_preproc_out mtmd_image_preprocessor_longest_edge::preprocess(const cl
const int cur_merge = hparams.n_merge == 0 ? 1 : hparams.n_merge;
const clip_image_size target_size = img_tool::calc_size_preserved_ratio(
original_size,
hparams.patch_size * cur_merge,
hparams.image_longest_edge);
{
/* align_size */ hparams.patch_size * cur_merge,
/* min_pixels */ std::max(0, hparams.image_min_pixels),
/* max_pixels */ std::max(0, hparams.image_max_pixels),
/* longest_edge */ hparams.image_longest_edge,
});
img_tool::resize(img, resized_image, target_size,
hparams.image_resize_algo,
hparams.image_resize_pad,
@@ -1000,8 +1003,8 @@ mtmd_image_preprocessor_llava_uhd::slice_instructions mtmd_image_preprocessor_lf
mtmd_image_preprocessor_llava_uhd::slice_instructions inst;
const int align_size = hparams.patch_size * hparams.n_merge;
inst.overview_size = img_tool::calc_size_preserved_ratio(
original_size, align_size,
hparams.image_min_pixels, hparams.image_max_pixels);
original_size,
{ align_size, hparams.image_min_pixels, hparams.image_max_pixels, 0 });
// tile if either dimension exceeds tile_size with tolerance
const bool needs_tiling = original_size.width > tile_size * max_pixels_tolerance || original_size.height > tile_size * max_pixels_tolerance;
@@ -1109,7 +1112,8 @@ mtmd_image_preproc_out mtmd_image_preprocessor_idefics3::preprocess(const clip_i
// CITE: https://github.com/huggingface/transformers/blob/main/src/transformers/models/idefics3/image_processing_idefics3.py#L737
const clip_image_size original_size = img.get_size();
const clip_image_size refined_size = img_tool::calc_size_preserved_ratio(
original_size, hparams.image_size, hparams.image_longest_edge);
original_size,
{ hparams.image_size, std::max(0, hparams.image_min_pixels), std::max(0, hparams.image_max_pixels), hparams.image_longest_edge });
// LOG_INF("%s: original size: %d x %d, refined size: %d x %d\n",
// __func__, original_size.width, original_size.height,
// refined_size.width, refined_size.height);
+248
View File
@@ -22,8 +22,123 @@
#include <cstdlib>
#include <cstring>
#include <climits>
#include <type_traits>
#include <vector>
// remember to bump this if the serialization format changes
#define MTMD_SERIALIZATION_VERSION 1
struct mtmd_serialization {
// note: using 64-bit here for future-proofing
uint64_t version = MTMD_SERIALIZATION_VERSION;
std::vector<char> data;
size_t read_pos = 0; // cursor used when reading
// for writing
mtmd_serialization(uint64_t version) : version(version) {
write(version);
}
// for reading
mtmd_serialization(uint64_t version, const char * buf, size_t len) {
// copy buf to data
data.assign(buf, buf + len);
uint64_t ver_in = read<uint64_t>();
if (ver_in != version) {
throw std::runtime_error("version mismatch");
}
this->version = ver_in;
}
template <typename T>
void write(T value) {
static_assert(std::is_trivially_copyable<T>::value && !std::is_same<T, bool>::value,
"T must be trivially copyable and not bool");
const char * p = reinterpret_cast<const char *>(&value);
data.insert(data.end(), p, p + sizeof(T));
}
template <typename T>
T read() {
static_assert(std::is_trivially_copyable<T>::value && !std::is_same<T, bool>::value,
"T must be trivially copyable and not bool");
if (read_pos + sizeof(T) > data.size()) {
throw std::runtime_error("read OOB");
}
T value;
std::memcpy(&value, data.data() + read_pos, sizeof(T));
read_pos += sizeof(T);
return value;
}
};
template <>
void mtmd_serialization::write<bool>(bool value) {
write<uint8_t>(value ? 1 : 0);
}
template <>
bool mtmd_serialization::read<bool>() {
return read<uint8_t>() != 0;
}
template <>
void mtmd_serialization::write<std::string>(std::string value) {
write<uint64_t>(value.size());
data.insert(data.end(), value.begin(), value.end());
}
template <>
std::string mtmd_serialization::read<std::string>() {
uint64_t len = read<uint64_t>();
if (read_pos + len > data.size()) {
throw std::runtime_error("read_string OOB");
}
std::string str(data.data() + read_pos, len);
read_pos += len;
return str;
}
// only mtmd.cpp needs these, so they're implemented here rather than in clip-impl.h
void clip_image_f32::serialize(mtmd_serialization & ser) const {
// remember to bump MTMD_SERIALIZATION_VERSION if this is changed
// note: buf is intentionally NOT serialized; the loaded clip_image_f32 will always be a placeholder
ser.write(add_viewsep);
ser.write(add_newline);
ser.write((int32_t)nx_);
ser.write((int32_t)ny_);
}
void clip_image_f32::deserialize(mtmd_serialization & ser) {
add_viewsep = ser.read<bool>();
add_newline = ser.read<bool>();
nx_ = ser.read<int32_t>();
ny_ = ser.read<int32_t>();
buf.clear(); // always a placeholder after loading
}
void clip_image_f32_batch::serialize(mtmd_serialization & ser) const {
// remember to bump MTMD_SERIALIZATION_VERSION if this is changed
ser.write(is_audio);
ser.write<uint64_t>(entries.size());
for (const auto & entry : entries) {
entry.serialize(ser);
}
}
void clip_image_f32_batch::deserialize(mtmd_serialization & ser) {
is_audio = ser.read<bool>();
uint64_t n = ser.read<uint64_t>();
constexpr size_t min_entry_bytes = sizeof(uint8_t) * 2 + sizeof(int32_t) * 2;
if (n > (ser.data.size() - ser.read_pos) / min_entry_bytes) {
throw std::runtime_error("entries count exceeds buffer size");
}
entries.clear();
entries.reserve(n);
for (uint64_t i = 0; i < n; i++) {
clip_image_f32 entry;
entry.deserialize(ser);
entries.push_back(std::move(entry));
}
}
// for still image data, layout is RGBRGBRGB...
// length of data must be nx * ny * 3 bytes
//
@@ -83,6 +198,7 @@ enum mtmd_pos_type {
MTMD_POS_TYPE_NORMAL, // number of positions equals to number of tokens
MTMD_POS_TYPE_MROPE, // qwen-vl mrope style, each image takes max(t,h,w) position indexes
MTMD_POS_TYPE_HUNYUANVL, // HunyuanVL mrope + BOI/EOI/newline layout with XD-RoPE dim-3
MTMD_POS_TYPE_COUNT, // for validation
};
struct mtmd_image_tokens {
@@ -136,6 +252,30 @@ struct mtmd_image_tokens {
id
};
}
void serialize(mtmd_serialization & ser) const {
// remember to bump MTMD_SERIALIZATION_VERSION if this is changed
ser.write(nx);
ser.write(ny);
ser.write((uint32_t)pos);
ser.write(image_idx);
ser.write(n_temporal_merge);
ser.write(id);
batch_f32.serialize(ser);
}
void deserialize(mtmd_serialization & ser) {
nx = ser.read<uint32_t>();
ny = ser.read<uint32_t>();
uint32_t pos_raw = ser.read<uint32_t>();
if (pos_raw >= MTMD_POS_TYPE_COUNT) {
throw std::runtime_error("invalid pos type");
}
pos = (mtmd_pos_type)pos_raw;
image_idx = ser.read<uint32_t>();
n_temporal_merge = ser.read<uint32_t>();
id = ser.read<std::string>();
batch_f32.deserialize(ser);
}
};
using mtmd_image_tokens_ptr = std::unique_ptr<mtmd_image_tokens>;
@@ -161,6 +301,18 @@ struct mtmd_audio_tokens {
id
};
}
void serialize(mtmd_serialization & ser) const {
// remember to bump MTMD_SERIALIZATION_VERSION if this is changed
ser.write(n_tokens);
ser.write(id);
batch_f32.serialize(ser);
}
void deserialize(mtmd_serialization & ser) {
n_tokens = ser.read<uint32_t>();
id = ser.read<std::string>();
batch_f32.deserialize(ser);
}
};
using mtmd_audio_tokens_ptr = std::unique_ptr<mtmd_audio_tokens>;
@@ -192,6 +344,66 @@ struct mtmd_input_chunk {
}
return false;
}
void serialize(mtmd_serialization & ser) const {
// remember to bump MTMD_SERIALIZATION_VERSION if this is changed
ser.write((uint32_t)type);
ser.write<uint64_t>(tokens_text.size());
for (llama_token tok : tokens_text) {
ser.write((int32_t)tok);
}
ser.write(tokens_image != nullptr);
if (tokens_image) {
tokens_image->serialize(ser);
}
ser.write(tokens_audio != nullptr);
if (tokens_audio) {
tokens_audio->serialize(ser);
}
}
void deserialize(mtmd_serialization & ser) {
uint32_t type_raw = ser.read<uint32_t>();
if (type_raw >= MTMD_INPUT_CHUNK_TYPE_COUNT) {
throw std::runtime_error("invalid chunk type");
}
type = (mtmd_input_chunk_type)type_raw;
uint64_t n_tokens_text = ser.read<uint64_t>();
// reject before resize() so a tiny corrupted/malicious buffer can't force a huge allocation
if (n_tokens_text > (ser.data.size() - ser.read_pos) / sizeof(int32_t)) {
throw std::runtime_error("tokens_text length exceeds buffer size");
}
tokens_text.resize(n_tokens_text);
for (uint64_t i = 0; i < n_tokens_text; i++) {
tokens_text[i] = (llama_token)ser.read<int32_t>();
}
if (ser.read<bool>()) {
tokens_image = std::make_unique<mtmd_image_tokens>();
tokens_image->deserialize(ser);
} else {
tokens_image.reset();
}
if (ser.read<bool>()) {
tokens_audio = std::make_unique<mtmd_audio_tokens>();
tokens_audio->deserialize(ser);
} else {
tokens_audio.reset();
}
// catch buffers where the declared type doesn't match which payload is actually present,
// so a mismatched chunk can't slip through and null-deref/abort later in an accessor
if (type == MTMD_INPUT_CHUNK_TYPE_IMAGE && !tokens_image) {
throw std::runtime_error("type is IMAGE but tokens_image is missing");
}
if (type == MTMD_INPUT_CHUNK_TYPE_AUDIO && !tokens_audio) {
throw std::runtime_error("type is AUDIO but tokens_audio is missing");
}
}
};
struct mtmd_input_chunks {
@@ -2043,6 +2255,42 @@ void mtmd_input_chunk_free(mtmd_input_chunk * chunk) {
}
}
int32_t mtmd_input_chunk_save(const mtmd_input_chunk * chunk, char * out_buf, size_t out_len, size_t * expected_out_len) {
try {
mtmd_serialization ser(MTMD_SERIALIZATION_VERSION);
chunk->serialize(ser);
if (expected_out_len) {
*expected_out_len = ser.data.size();
}
if (!out_buf) {
// caller is only querying the required size
return 0;
}
if (out_len < ser.data.size()) {
LOG_ERR("%s: out_buf is too small, need %zu bytes, got %zu\n", __func__, ser.data.size(), out_len);
return -1;
}
std::memcpy(out_buf, ser.data.data(), ser.data.size());
return 0;
} catch (const std::exception & e) {
LOG_ERR("%s: %s\n", __func__, e.what());
return -1;
}
}
mtmd_input_chunk * mtmd_input_chunk_load(const char * buf, size_t len) {
try {
mtmd_serialization ser(MTMD_SERIALIZATION_VERSION, buf, len);
mtmd::input_chunk_ptr chunk(new mtmd_input_chunk());
chunk->deserialize(ser);
return chunk.release();
} catch (const std::exception & e) {
LOG_ERR("%s: %s\n", __func__, e.what());
return nullptr;
}
}
// mtmd_image_tokens
size_t mtmd_image_tokens_get_n_tokens(const mtmd_image_tokens * image_tokens) {
+10
View File
@@ -55,6 +55,7 @@ enum mtmd_input_chunk_type {
MTMD_INPUT_CHUNK_TYPE_TEXT,
MTMD_INPUT_CHUNK_TYPE_IMAGE,
MTMD_INPUT_CHUNK_TYPE_AUDIO,
MTMD_INPUT_CHUNK_TYPE_COUNT, // for validation
};
// opaque types
@@ -232,6 +233,15 @@ MTMD_API llama_pos mtmd_input_chunk_get_n_pos (const mtmd
MTMD_API mtmd_input_chunk * mtmd_input_chunk_copy(const mtmd_input_chunk * chunk);
MTMD_API void mtmd_input_chunk_free(mtmd_input_chunk * chunk);
// save/load an input chunk to/from a buffer (useful for KV save/load)
// important: only chunk's metadata will be saved, the actual image/audio data will not be saved
// the loaded chunk will always be a placeholder, cannot be used for mtmd_encode() or mtmd_batch_encode()
// out_buf can be nullptr (to query expected_out_len)
// returns 0 on success, non-zero on failure
MTMD_API int32_t mtmd_input_chunk_save(const mtmd_input_chunk * chunk, char * out_buf, size_t out_len, size_t * expected_out_len);
// returns nullptr on failure
MTMD_API mtmd_input_chunk * mtmd_input_chunk_load(const char * buf, size_t len);
// mtmd_image_tokens
//
+352 -44
View File
@@ -70,6 +70,188 @@ struct server_subproc {
}
};
struct server_lru_sched {
server_lru_sched(server_models & models) : models(models) {}
bool has_capacity(std::unique_lock<std::mutex> & lk) {
check_lock(lk);
return models.base_params.models_max <= 0
|| count_running() < (size_t) models.base_params.models_max;
}
// returns "" if no model can be given up
std::string pick_victim(std::unique_lock<std::mutex> & lk, const std::string & exclude) {
check_lock(lk);
std::string victim;
int64_t victim_last_used = 0;
for (const auto & m : models.mapping) {
if (m.first == exclude) {
continue;
}
// a busy model is mid-request, one still coming up has no request to finish
if (m.second.req_count != 0 || !m.second.meta.is_ready_or_sleep()) {
continue;
}
if (victim.empty() || m.second.meta.last_used < victim_last_used) {
victim = m.first;
victim_last_used = m.second.meta.last_used;
}
}
return victim;
}
// requests wanting the same model share one entry, so they all need only one slot
// and all get unblocked by the single load that entry performs
void join(std::unique_lock<std::mutex> & lk, const std::string & model_id) {
check_lock(lk);
if (entry_t * e = find(model_id)) {
e->n_waiters++;
SRV_INF("request for name=%s joined the queue, %d waiting\n", model_id.c_str(), e->n_waiters);
return;
}
queue.push_back({ model_id, 1, false, false });
SRV_INF("models_max reached, request for name=%s queued at position %zu\n",
model_id.c_str(), queue.size());
}
void leave(std::unique_lock<std::mutex> & lk, const std::string & model_id) {
check_lock(lk);
for (auto it = queue.begin(); it != queue.end(); ++it) {
if (it->model_id == model_id) {
if (--it->n_waiters <= 0) {
queue.erase(it); // last one waiting for this model went away
}
return;
}
}
}
bool queue_empty(std::unique_lock<std::mutex> & lk) {
check_lock(lk);
return queue.empty();
}
// true if it is this model's turn to load, and nobody is loading it yet
bool try_claim(std::unique_lock<std::mutex> & lk, const std::string & model_id) {
check_lock(lk);
if (queue.empty() || queue.front().model_id != model_id || queue.front().loading) {
return false;
}
if (!has_capacity(lk)) {
return false;
}
queue.front().loading = true;
return true;
}
// ok means the model is up: drop the entry, the other waiters just watch its status now
void claim_done(std::unique_lock<std::mutex> & lk, const std::string & model_id, bool ok) {
check_lock(lk);
for (auto it = queue.begin(); it != queue.end(); ++it) {
if (it->model_id == model_id) {
if (ok) {
queue.erase(it);
} else {
it->loading = false;
}
return;
}
}
}
// a model is on its way out for this entry, so other requests do not also give up one
void mark_slot_pending(std::unique_lock<std::mutex> & lk, const std::string & model_id) {
check_lock(lk);
if (entry_t * e = find(model_id)) {
e->slot_pending = true;
}
}
// model_id went idle: give up its slot if a queued request needs one
// thread-safe, caller must NOT hold models.mutex
void on_model_idle(const std::string & model_id) {
if (models.base_params.models_max <= 0) {
return; // no limit, nothing is ever queued
}
{
std::unique_lock<std::mutex> lk(models.mutex);
if (queue.empty()) {
return;
}
size_t promised = 0;
bool has_unserved = false;
for (const auto & e : queue) {
if (e.needs_slot()) {
has_unserved = true;
} else {
promised++;
}
}
if (!has_unserved) {
return;
}
if ((int) count_running() - (int) promised < models.base_params.models_max) {
return; // a slot is already on its way
}
// never give up a model that a queued request wants
for (const auto & e : queue) {
if (e.model_id == model_id) {
return;
}
}
auto it = models.mapping.find(model_id);
if (it == models.mapping.end() || it->second.req_count != 0 || !it->second.meta.is_ready_or_sleep()) {
return;
}
for (auto & e : queue) {
if (!e.slot_pending) {
e.slot_pending = true;
break;
}
}
}
SRV_INF("model name=%s went idle, giving up its slot to a queued request\n", model_id.c_str());
models.unload(model_id);
}
private:
struct entry_t {
std::string model_id;
int n_waiters; // requests waiting for this model
bool slot_pending; // a model is already being evicted for this entry
bool loading; // one of the waiters is doing the load right now
// a slot is already coming, or already taken by the load in flight
bool needs_slot() const { return !slot_pending && !loading; }
};
entry_t * find(const std::string & model_id) {
for (auto & e : queue) {
if (e.model_id == model_id) {
return &e;
}
}
return nullptr;
}
void check_lock(std::unique_lock<std::mutex> & lk) {
GGML_ASSERT(lk.owns_lock() && lk.mutex() == &models.mutex);
}
size_t count_running() {
size_t count = 0;
for (const auto & m : models.mapping) {
if (m.second.meta.is_running()) {
count++;
}
}
return count;
}
server_models & models;
std::deque<entry_t> queue;
};
// short loopback budget for the resumable stream router to child JSON calls (probe, lookup,
// delete). distinct from params.timeout_read/write which only applies to the generation proxy
static constexpr int STREAM_LOOKUP_TIMEOUT_MS = 250;
@@ -229,7 +411,8 @@ server_models::server_models(
: ctx_preset(LLAMA_EXAMPLE_SERVER),
base_params(params),
base_env(get_environment()),
base_preset(ctx_preset.load_from_args(argc, argv)) {
base_preset(ctx_preset.load_from_args(argc, argv)),
sched(std::make_unique<server_lru_sched>(*this)) {
// clean up base preset
unset_reserved_args(base_preset, true);
// set binary path
@@ -241,8 +424,11 @@ server_models::server_models(
LOG_WRN("using original argv[0] as fallback: %s\n", argv[0]);
}
load_models();
debug_fake_timing = !common_get_env("LLAMA_SERVER_DEBUG_FAKE_TIMING").empty();
}
server_models::~server_models() = default;
void server_models::add_model(server_model_meta && meta) {
if (mapping.find(meta.name) != mapping.end()) {
throw std::runtime_error(string_format("model '%s' appears multiple times", meta.name.c_str()));
@@ -713,22 +899,15 @@ void server_models::unload_lru() {
return; // no limit
}
// remove one of the servers if we passed the models_max (least recently used - LRU)
std::string lru_model_name = "";
int64_t lru_last_used = ggml_time_ms();
size_t count_active = 0;
std::string lru_model_name;
{
std::unique_lock<std::mutex> lk(mutex);
for (const auto & m : mapping) {
if (m.second.meta.is_running()) {
count_active++;
if (m.second.meta.last_used < lru_last_used) {
lru_model_name = m.first;
lru_last_used = m.second.meta.last_used;
}
}
if (sched->has_capacity(lk)) {
return;
}
lru_model_name = sched->pick_victim(lk, "");
}
if (!lru_model_name.empty() && count_active >= (size_t)base_params.models_max) {
if (!lru_model_name.empty()) {
SRV_INF("models_max limit reached, removing LRU name=%s\n", lru_model_name.c_str());
unload(lru_model_name);
// wait for unload to complete
@@ -746,6 +925,11 @@ void server_models::load(const std::string & name) {
}
void server_models::load(const std::string & name, const load_options & opts) {
if (debug_fake_timing) {
// do not hold the mutex here, other requests must keep making progress
std::this_thread::sleep_for(std::chrono::seconds(2));
}
if (!opts.custom_meta.has_value()) {
if (!has_model(name)) {
throw std::runtime_error("model name=" + name + " is not found");
@@ -1138,7 +1322,7 @@ void server_models::wait(std::unique_lock<std::mutex> & lk, const std::string &
});
}
bool server_models::ensure_model_ready(const std::string & name) {
bool server_models::ensure_model_ready(const std::string & name, const std::function<bool()> & should_stop) {
auto meta = get_meta(name);
if (!meta.has_value()) {
throw std::runtime_error("model name=" + name + " is not found");
@@ -1149,25 +1333,112 @@ bool server_models::ensure_model_ready(const std::string & name) {
if (meta->status == SERVER_MODEL_STATUS_SLEEPING) {
return false; // child is sleeping but still running; new request will wake it up
}
if (meta->status == SERVER_MODEL_STATUS_UNLOADED) {
SRV_INF("model name=%s is not loaded, loading...\n", name.c_str());
load(name);
}
// wait for loading to complete
SRV_INF("waiting until model name=%s is fully loaded...\n", name.c_str());
wait(name, [&meta](const server_model_meta & new_meta) {
if (new_meta.status != SERVER_MODEL_STATUS_LOADING) {
meta = new_meta; // update meta for final check after wait
return true;
bool queued = false;
bool did_load = false;
std::string victim;
{
std::unique_lock<std::mutex> lk(mutex);
auto it = mapping.find(name);
if (it != mapping.end() && it->second.meta.status == SERVER_MODEL_STATUS_UNLOADED) {
bool has_capacity = sched->has_capacity(lk);
if (has_capacity && sched->queue_empty(lk)) {
lk.unlock();
SRV_INF("model name=%s is not loaded, loading...\n", name.c_str());
load(name);
did_load = true;
} else {
// also queue when a slot looks free but others wait already, else they starve
sched->join(lk, name);
queued = true;
if (!has_capacity) {
// an idle model may sit here right now, do not wait for a request to end
victim = sched->pick_victim(lk, name);
if (!victim.empty()) {
sched->mark_slot_pending(lk, name);
}
}
}
}
return false;
});
// check final status
if (!meta.has_value() || meta->is_failed()) {
throw std::runtime_error("model name=" + name + " failed to load");
}
if (!victim.empty()) {
SRV_INF("evicting idle LRU name=%s to make room for name=%s\n", victim.c_str(), name.c_str());
unload(victim);
}
// while queued, this is also where the load happens: the head of the queue does it
SRV_INF("waiting until model name=%s is fully loaded...\n", name.c_str());
std::unique_lock<std::mutex> lk(mutex);
auto leave_queue = [this, &queued, &lk, &name]() {
if (queued) {
sched->leave(lk, name);
queued = false;
}
};
try {
bool saw_loading = false;
while (true) {
auto it = mapping.find(name);
if (it == mapping.end()) {
break; // removed by another code path, nothing to wait for
}
const server_model_status status = it->second.meta.status;
if (status == SERVER_MODEL_STATUS_LOADED || status == SERVER_MODEL_STATUS_SLEEPING) {
break;
}
if (status == SERVER_MODEL_STATUS_DOWNLOADING || status == SERVER_MODEL_STATUS_DOWNLOADED) {
break; // do not wait on a download child
}
if (status == SERVER_MODEL_STATUS_LOADING) {
saw_loading = true;
} else if (status == SERVER_MODEL_STATUS_UNLOADED) {
if (did_load || saw_loading) {
// a spawn happened and the instance came back down
if (it->second.meta.is_failed()) {
throw std::runtime_error("model name=" + name + " failed to load");
}
break; // unloaded by another code path, caller reports "not running"
}
if (!queued) {
break; // not queued, and the load someone else started fell over
}
}
if (should_stop && should_stop()) {
// if a model was evicted for us, the free slot goes to the next waiter
throw std::runtime_error("request cancelled while waiting for model name=" + name);
}
// our turn: our model is at the head, and a slot really did free up
if (status == SERVER_MODEL_STATUS_UNLOADED && sched->try_claim(lk, name)) {
lk.unlock();
bool ok = true;
try {
SRV_INF("slot available, loading queued model name=%s\n", name.c_str());
load(name);
did_load = true;
} catch (const std::exception & e) {
// lost a race for the slot, stay in line and retry
SRV_WRN("queued load of name=%s did not go through: %s\n", name.c_str(), e.what());
ok = false;
}
lk.lock();
sched->claim_done(lk, name, ok);
if (ok) {
queued = false; // entry is gone, the other waiters watch the status now
}
continue;
}
cv.wait_for(lk, std::chrono::milliseconds(200));
}
} catch (...) {
leave_queue();
throw;
}
leave_queue();
return true;
}
@@ -1180,9 +1451,16 @@ server_http_res_ptr server_models::proxy_request(const server_http_req & req, co
if (!meta->is_running()) {
throw std::invalid_argument("model name=" + name + " is not running");
}
if (update_last_used) {
{
std::unique_lock<std::mutex> lk(mutex);
mapping[name].meta.last_used = ggml_time_ms();
if (update_last_used) {
mapping[name].meta.last_used = ggml_time_ms();
}
mapping[name].req_count++;
}
if (debug_fake_timing) {
// sleep after req_count++, so the model counts as busy while we wait here
std::this_thread::sleep_for(std::chrono::seconds(2));
}
SRV_INF("proxying request to model %s on port %d\n", name.c_str(), meta->port);
std::string proxy_path = req.path;
@@ -1198,13 +1476,29 @@ server_http_res_ptr server_models::proxy_request(const server_http_req & req, co
req.headers,
req.body,
req.files,
// a detached request belongs to a replay session that outlives the client socket:
// it reaches the child even when the downstream died during the load wait, the
// session buffer is the recipient and DELETE remains the stop
detached ? std::function<bool()>([]() { return false; }) : req.should_stop,
// a detached request belongs to a replay session
detached
? std::function<bool()>([]() { return false; })
: req.should_stop,
base_params.timeout_read,
base_params.timeout_write
);
proxy->cleanup = [this, name]() {
bool went_idle = false;
{
std::unique_lock<std::mutex> lk(mutex);
auto it = mapping.find(name);
if (it != mapping.end() && it->second.req_count > 0) {
it->second.req_count--;
went_idle = it->second.req_count == 0;
}
}
if (went_idle) {
sched->on_model_idle(name);
}
};
return proxy;
}
@@ -1568,7 +1862,7 @@ void server_models_routes::init_routes() {
return error_res;
}
if (autoload) {
models.ensure_model_ready(name);
models.ensure_model_ready(name, req.should_stop);
}
return models.proxy_request(req, method, name, false);
};
@@ -1588,7 +1882,9 @@ void server_models_routes::init_routes() {
// this request instead of leaving an orphan generation
std::string conv_id = server_stream_conv_id_from_headers(req.headers);
uint64_t ticket = models.conv_models.remember(conv_id, name);
bool waited = autoload && models.ensure_model_ready(name);
// a dead socket must not cancel a session request, only a stop does (checked right below)
auto should_stop = ticket == 0 ? req.should_stop : nullptr;
bool waited = autoload && models.ensure_model_ready(name, should_stop);
if (ticket != 0 && !models.conv_models.alive(conv_id, ticket)) {
SRV_INF("request for conv_id=%s cancelled while model name=%s was loading\n",
conv_id.c_str(), name.c_str());
@@ -2064,7 +2360,7 @@ server_http_proxy::server_http_proxy(
cli->set_write_timeout(timeout_read, 0); // reversed for cli (client) vs srv (server)
cli->set_read_timeout(timeout_write, 0);
this->status = 500; // to be overwritten upon response
this->cleanup = [pipe]() {
this->cleanup_pipes = [pipe]() {
pipe->close_read();
pipe->close_write();
};
@@ -2079,9 +2375,8 @@ server_http_proxy::server_http_proxy(
return has_next; // false if EOF or pipe broken
};
// wire up the HTTP client
// note: do NOT capture `this` pointer, as it may be destroyed before the thread ends
httplib::ResponseHandler response_handler = [pipe, cli](const httplib::Response & response) {
// build the header message forwarded to the reader thread, stripping internal proxy headers
auto make_header_msg = [](const httplib::Response & response) {
msg_t msg;
msg.status = response.status;
for (const auto & [key, value] : response.headers) {
@@ -2095,7 +2390,17 @@ server_http_proxy::server_http_proxy(
}
msg.headers[key] = value;
}
return pipe->write(std::move(msg)); // send headers first
return msg;
};
// true once response_handler has already forwarded the headers
auto headers_sent = std::make_shared<std::atomic<bool>>(false);
// wire up the HTTP client
// note: do NOT capture `this` pointer, as it may be destroyed before the thread ends
httplib::ResponseHandler response_handler = [pipe, headers_sent, make_header_msg](const httplib::Response & response) {
headers_sent->store(true);
return pipe->write(make_header_msg(response)); // send headers first
};
httplib::ContentReceiverWithProgress content_receiver = [pipe](const char * data, size_t data_length, size_t, size_t) {
// send data chunks
@@ -2169,13 +2474,16 @@ server_http_proxy::server_http_proxy(
// start the proxy thread
SRV_DBG("start proxy thread %s %s\n", req.method.c_str(), req.path.c_str());
this->thread = std::thread([cli, pipe, req]() {
this->thread = std::thread([cli, pipe, req, headers_sent, make_header_msg]() {
auto result = cli->send(std::move(req));
if (result.error() != httplib::Error::Success) {
auto err_str = httplib::to_string(result.error());
SRV_ERR("http client error: %s\n", err_str.c_str());
pipe->write({{}, 500, "", ""}); // header
pipe->write({{}, 0, "proxy error: " + err_str, ""}); // body
} else if (!headers_sent->load()) {
// httplib skips response_handler for bodyless statuses like 204, send headers here instead
pipe->write(make_header_msg(*result));
}
pipe->close_write(); // signal EOF to reader
SRV_DBG("%s", "client request thread ended\n");
+22 -4
View File
@@ -84,7 +84,6 @@ struct server_model_meta {
int exit_code = 0; // exit code of the model instance process (only valid if status == FAILED)
int stop_timeout = 0; // seconds to wait before force-killing the model instance during shutdown
mtmd_caps multimodal; // multimodal capabilities
// bool need_download = false; // whether the model needs to be downloaded before loading // TODO @ngxson: implement this
bool is_ready() const {
return status == SERVER_MODEL_STATUS_LOADED;
@@ -94,6 +93,10 @@ struct server_model_meta {
return status == SERVER_MODEL_STATUS_LOADED || status == SERVER_MODEL_STATUS_LOADING || status == SERVER_MODEL_STATUS_SLEEPING;
}
bool is_ready_or_sleep() const {
return status == SERVER_MODEL_STATUS_LOADED || status == SERVER_MODEL_STATUS_SLEEPING;
}
bool is_failed() const {
return status == SERVER_MODEL_STATUS_UNLOADED && exit_code != 0;
}
@@ -103,16 +106,19 @@ struct server_model_meta {
};
struct server_models_routes;
struct server_subproc; // defined in server-models.cpp
struct server_subproc; // defined in server-models.cpp
struct server_lru_sched; // defined in server-models.cpp
struct server_models {
friend struct server_models_routes;
friend struct server_lru_sched;
private:
struct instance_t {
std::shared_ptr<server_subproc> subproc; // shared between main thread and monitoring thread
std::thread th;
server_model_meta meta;
int req_count = 0; // number of active proxy requests
};
std::mutex mutex;
@@ -191,6 +197,12 @@ private:
std::vector<std::string> base_env;
common_preset base_preset; // base preset from llama-server CLI args
// queue of requests waiting for a models_max slot
std::unique_ptr<server_lru_sched> sched;
// if true, add some delay to simulate works (useful for testing)
bool debug_fake_timing = false;
void update_meta(const std::string & name, const server_model_meta & meta);
// unload least recently used models if the limit is reached
@@ -207,6 +219,7 @@ public:
conv_model_tracker conv_models;
server_models(const common_params & params, int argc, char ** argv);
~server_models();
server_response sse; // for real-time updates via SSE endpoint
@@ -263,7 +276,9 @@ public:
// ensure the model is in ready state (thread-safe)
// return false if model is ready
// otherwise, load the model and blocking wait until it's ready, then return true (meta may need to be refreshed)
bool ensure_model_ready(const std::string & name);
// if models_max is reached, the request waits in a queue until a slot frees up
// throws if the load fails, or if should_stop fires while waiting
bool ensure_model_ready(const std::string & name, const std::function<bool()> & should_stop = nullptr);
// proxy an HTTP request to the model instance
server_http_res_ptr proxy_request(const server_http_req & req, const std::string & method, const std::string & name, bool update_last_used, bool detached = false);
@@ -343,7 +358,6 @@ struct server_models_routes {
*/
struct server_http_proxy : server_http_res {
std::function<void()> cleanup = nullptr;
public:
server_http_proxy(const std::string & method,
const std::string & scheme,
const std::string & host,
@@ -357,11 +371,15 @@ public:
int32_t timeout_write
);
~server_http_proxy() {
if (cleanup_pipes) {
cleanup_pipes();
}
if (cleanup) {
cleanup();
}
}
private:
std::function<void()> cleanup_pipes = nullptr;
std::thread thread;
struct msg_t {
std::map<std::string, std::string> headers;
+1 -1
View File
@@ -519,7 +519,7 @@ task_params eval_llama_cmpl_schema(
const json & data) {
task_params params;
// Sampling parameter defaults are loaded from the global server context (but individual requests can still them)
// Sampling parameter defaults are loaded from the global server context (but individual requests can still override them)
params.sampling = params_base.sampling;
params.speculative = params_base.speculative;
params.n_keep = params_base.n_keep;
+30
View File
@@ -1,5 +1,7 @@
import pytest
from utils import *
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
server = ServerPreset.tinyllama2()
@@ -39,3 +41,31 @@ def test_mcp_proxy_custom_port():
res = server.make_request("GET", f"/cors-proxy?url=http://{server.server_host}:{server.server_port}/models")
assert res.status_code == 200
assert "data" in res.body
def test_mcp_proxy_no_content():
# note: see issue #26598
class NoContentHandler(BaseHTTPRequestHandler):
def do_POST(self):
self.send_response(204)
self.end_headers()
def log_message(self, format, *args):
pass
target = ThreadingHTTPServer(("127.0.0.1", 0), NoContentHandler)
target_thread = threading.Thread(target=target.serve_forever, daemon=True)
target_thread.start()
try:
global server
server.ui_mcp_proxy = True
server.start()
res = server.make_request("POST", f"/cors-proxy?url=http://127.0.0.1:{target.server_port}/", data={})
assert res.status_code == 204
assert res.body in (None, b"", "")
finally:
target.shutdown()
target.server_close()
+150
View File
@@ -145,6 +145,156 @@ def test_router_models_max_evicts_lru():
assert _get_model_status(first) == "unloaded"
# server_lru_sched tests (relying on LLAMA_SERVER_DEBUG_FAKE_TIMING)
MODEL_A = "ggml-org/tinygemma3-GGUF:Q8_0"
MODEL_B = "ggml-org/test-model-stories260K:F32"
MODEL_C = "ggml-org/test-model-stories260K-infill:F32"
def _tokenize(model_id: str, timeout: float | None = DEFAULT_REQUEST_TIMEOUT) -> ServerResponse:
return server.make_request(
"POST", "/tokenize", data={"model": model_id, "content": "hello world"}, timeout=timeout
)
class _Bg:
"""runs one request in a thread, keeps its result, error and finish time"""
def __init__(self, fn):
self.result = None
self.error: Exception | None = None
self.done_at: float = 0.0
self._thread = threading.Thread(target=self._run, args=(fn,), daemon=True)
def _run(self, fn):
try:
self.result = fn()
except Exception as e:
self.error = e
self.done_at = time.time()
def start(self):
self._thread.start()
return self
def join(self, timeout: int = 180):
self._thread.join(timeout)
assert not self._thread.is_alive(), "background request did not finish in time"
return self
def assert_ok(self, what: str):
assert self.error is None, f"{what} raised {self.error!r}"
assert self.result is not None and self.result.status_code == 200, \
f"{what} failed: {self.result.status_code if self.result else None} {self.result.body if self.result else None}"
def test_router_queue_does_not_evict_busy_model():
"""a request that finds no free slot waits, and the model serving a request survives it"""
global server
server.models_max = 1
server.start()
_load_model_and_wait(MODEL_A, timeout=120)
busy = _Bg(lambda: _tokenize(MODEL_A)).start()
time.sleep(0.5) # let the request reach the child and take the only slot
# no slot free and MODEL_A is busy, so this queues instead of evicting mid-request
queued = _Bg(lambda: _tokenize(MODEL_B)).start()
busy.join()
queued.join()
# had MODEL_A been evicted while serving, its own request would have died
busy.assert_ok("request against the busy model")
queued.assert_ok("queued request")
_wait_for_model_status(MODEL_B, {"loaded"}, timeout=120)
assert _get_model_status(MODEL_A) == "unloaded"
def test_router_queue_coalesces_requests_for_same_model():
"""many requests for one missing model share a slot, so only one model is given up"""
global server
server.models_max = 2
server.start()
_load_model_and_wait(MODEL_A, timeout=120)
_load_model_and_wait(MODEL_B, timeout=120)
# keep MODEL_A busy so MODEL_B is the only model that can be given up
busy = _Bg(lambda: _tokenize(MODEL_A)).start()
time.sleep(0.5)
waiters = [_Bg(lambda: _tokenize(MODEL_C)).start() for _ in range(3)]
busy.join()
for w in waiters:
w.join()
busy.assert_ok("request against the busy model")
for i, w in enumerate(waiters):
w.assert_ok(f"queued request {i}")
_wait_for_model_status(MODEL_C, {"loaded"}, timeout=120)
# one entry for 3 requests means one eviction: MODEL_B goes, MODEL_A is left alone.
# without coalescing the leftover entries still ask for a slot,
# and MODEL_A is taken too as soon as it goes idle
assert _get_model_status(MODEL_A) == "loaded"
assert _get_model_status(MODEL_B) == "unloaded"
def test_router_queue_client_disconnect_keeps_model():
"""a client that leaves while queued must not cost a running model its slot"""
global server
server.models_max = 1
server.start()
_load_model_and_wait(MODEL_A, timeout=120)
busy = _Bg(lambda: _tokenize(MODEL_A)).start()
time.sleep(0.5)
# queues behind MODEL_A, then gives up long before MODEL_A goes idle
with pytest.raises(requests.exceptions.RequestException):
_tokenize(MODEL_B, timeout=1)
busy.join()
busy.assert_ok("request against the busy model")
# nobody is waiting anymore, so MODEL_A keeps its slot
time.sleep(3)
assert _get_model_status(MODEL_A) == "loaded"
assert _get_model_status(MODEL_B) == "unloaded"
def test_router_queue_is_fifo():
"""the queue is served in arrival order"""
global server
server.models_max = 1
server.start()
_load_model_and_wait(MODEL_A, timeout=120)
busy = _Bg(lambda: _tokenize(MODEL_A)).start()
time.sleep(0.5)
first = _Bg(lambda: _tokenize(MODEL_B)).start()
time.sleep(1) # keep the arrival order unambiguous
second = _Bg(lambda: _tokenize(MODEL_C)).start()
busy.join()
first.join()
second.join()
busy.assert_ok("request against the busy model")
first.assert_ok("first queued request")
second.assert_ok("second queued request")
assert first.done_at < second.done_at, "queue was not served in arrival order"
def test_router_no_models_autoload():
global server
server.no_models_autoload = True
+4 -1
View File
@@ -132,7 +132,10 @@ class ServerProcess:
self.external_server = "DEBUG_EXTERNAL" in os.environ
def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None:
env = {**os.environ}
env = {
**os.environ,
"LLAMA_SERVER_DEBUG_FAKE_TIMING": "1",
}
if "LLAMA_CACHE" not in os.environ:
env["LLAMA_CACHE"] = "tmp"
if self.external_server:
+1
View File
@@ -1,2 +1,3 @@
engine-strict=true
ignore-scripts=true
min-release-age=7
+2 -4
View File
@@ -143,10 +143,8 @@ declare global {
idxThemeStyle?: number;
idxCodeBlock?: number;
// File System Access API - missing from older DOM lib versions.
// Used by ChatFormWorkingDirectory's native folder picker. Feature availability
// is gated at runtime via `typeof window.showDirectoryPicker === 'function'`.
showDirectoryPicker: (options?: {
// File System Access API - not in the DOM lib and unavailable in some browsers
showDirectoryPicker?: (options?: {
id?: string;
mode?: 'read' | 'readwrite';
startIn?: FileSystemHandle | string;
@@ -14,9 +14,7 @@
INPUT_CLASSES,
SETTING_CONFIG_DEFAULT,
INITIAL_FILE_SIZE,
PROMPT_CONTENT_SEPARATOR,
PROMPT_TRIGGER_PREFIX,
RESOURCE_TRIGGER_PREFIX
PROMPT_CONTENT_SEPARATOR
} from '$lib/constants';
import {
ContentPartType,
@@ -39,8 +37,23 @@
activeConversation,
pendingCwd
} from '$lib/stores/conversations.svelte';
import type { GetPromptResult, MCPPromptInfo, MCPResourceInfo, PromptMessage } from '$lib/types';
import { isIMEComposing, parseClipboardContent, uuid } from '$lib/utils';
import type {
FileMentionEntry,
GetPromptResult,
MCPPromptInfo,
MCPResourceInfo,
PromptMessage
} from '$lib/types';
import {
buildMentionInsertion,
findCommandToken,
findMentionToken,
isIMEComposing,
mentionLinkEndingAt,
parseClipboardContent,
uuid
} from '$lib/utils';
import { useChatFormPickers } from '$lib/hooks/use-chat-form-pickers.svelte';
import {
AudioRecorder,
convertToWav,
@@ -96,30 +109,52 @@
onValueChange
}: Props = $props();
// Component References
let audioRecorder: AudioRecorder | undefined;
let chatFormActionsRef: ChatFormActions | undefined = $state(undefined);
let fileInputRef: ChatFormFileInputInvisible | undefined = $state(undefined);
let pickersRef: { handleKeydown: (event: KeyboardEvent) => boolean } | undefined =
$state(undefined);
let textareaRef: ChatFormTextarea | undefined = $state(undefined);
let inputRef: ChatFormTextarea | undefined = $state(undefined);
// Audio Recording State
let isRecording = $state(false);
let recordingSupported = $state(false);
// Picker State
let isPromptPickerOpen = $state(false);
let promptSearchQuery = $state('');
let isInlineResourcePickerOpen = $state(false);
let resourceSearchQuery = $state('');
// Invisible anchor at the form's top edge so the mention/WD popovers
// float above the box.
let mentionAnchor: HTMLDivElement | null = $state(null);
let cwd = $derived(activeConversation()?.cwd ?? pendingCwd());
async function handleWorkingDirectoryChange(value: string | null) {
await conversationsStore.setCwd(value);
const pickers = useChatFormPickers({
getValue: () => value,
setValue: (v) => {
value = v;
onValueChange?.(v);
},
getCaretOffset: () => inputRef?.getCaretOffset(),
setCaretOffset: (offset) => inputRef?.setCaretOffset(offset),
focusInput: refocusInput,
getShowModelSelector: () => showModelSelector,
hasPrompts: () => mcpStore.hasPromptsCapability(conversationsStore.getAllMcpServerOverrides()),
hasBuiltinTools: () => toolsStore.builtinTools.length > 0,
getCwd: () => cwd,
getServerHome: () => toolsStore.serverHome ?? null,
openModelSelector: () => chatFormActionsRef?.openModelSelector(),
getPickersRef: () => pickersRef
});
async function handleWorkingDirectoryChange(newDir: string | null) {
// Committing a directory consumes the `/cwd` token; the chip's
// clear-X path has no token to consume.
const token = findCommandToken(value);
if (token && token.name === 'cwd') {
value = '';
onValueChange?.('');
}
await conversationsStore.setCwd(newDir);
if (conversationsStore.activeConversation) {
await chatStore.recordCwdChange(value?.trim() || null);
await chatStore.recordCwdChange(newDir?.trim() || null);
}
}
@@ -171,18 +206,12 @@
audioRecorder = new AudioRecorder();
});
// Defer so the closing popover's focus scope tears down first - bits-ui
// yanks a synchronous focus() back into the still-mounted popover.
function refocusInput() {
queueMicrotask(() => textareaRef?.focus());
}
export function focus() {
textareaRef?.focus();
inputRef?.focus();
}
export function resetTextareaHeight() {
textareaRef?.resetHeight();
inputRef?.resetHeight();
}
export function openModelSelector() {
@@ -216,47 +245,26 @@
}
}
function handleInput() {
const perChatOverrides = conversationsStore.getAllMcpServerOverrides();
const hasServers = mcpStore.hasEnabledServers(perChatOverrides);
if (value.startsWith(PROMPT_TRIGGER_PREFIX) && hasServers) {
isPromptPickerOpen = true;
promptSearchQuery = value.slice(1);
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
} else if (
value.startsWith(RESOURCE_TRIGGER_PREFIX) &&
hasServers &&
mcpStore.hasResourcesCapability(perChatOverrides)
) {
isInlineResourcePickerOpen = true;
resourceSearchQuery = value.slice(1);
isPromptPickerOpen = false;
promptSearchQuery = '';
} else {
isPromptPickerOpen = false;
promptSearchQuery = '';
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
}
}
function handleKeydown(event: KeyboardEvent) {
if (pickersRef?.handleKeydown(event)) {
// Pickers consume navigation/escape keys first; when consumed, skip
// the enter-to-submit logic below.
if (pickers.handleKeydown(event)) {
return;
}
if (event.key === KeyboardKey.ESCAPE && isPromptPickerOpen) {
isPromptPickerOpen = false;
promptSearchQuery = '';
return;
}
if (event.key === KeyboardKey.ESCAPE && isInlineResourcePickerOpen) {
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
return;
// Backspace at a mention link's end deletes the whole token at once.
if (event.key === KeyboardKey.BACKSPACE && !event.ctrlKey && !event.metaKey && !event.altKey) {
const el = inputRef?.getElement();
if (el instanceof HTMLTextAreaElement && el.selectionStart === el.selectionEnd) {
const link = mentionLinkEndingAt(value, el.selectionStart);
if (link) {
event.preventDefault();
value = value.slice(0, link.start) + value.slice(link.end);
onValueChange?.(value);
queueMicrotask(() => inputRef?.setCaretOffset(link.start));
return;
}
}
}
if (event.key === KeyboardKey.ENTER && !event.shiftKey && !isIMEComposing(event)) {
@@ -332,7 +340,7 @@
}
setTimeout(() => {
textareaRef?.focus();
inputRef?.focus();
}, 10);
return;
@@ -359,13 +367,7 @@
promptInfo: MCPPromptInfo,
args?: Record<string, string>
) {
// Only clear the value if the prompt was triggered by typing '/'
if (value.startsWith(PROMPT_TRIGGER_PREFIX)) {
value = '';
onValueChange?.('');
}
isPromptPickerOpen = false;
promptSearchQuery = '';
pickers.closePromptPicker();
const promptName = promptInfo.title || promptInfo.name;
const placeholder: ChatUploadedFile = {
@@ -384,7 +386,7 @@
uploadedFiles = [...uploadedFiles, placeholder];
onUploadedFilesChange?.(uploadedFiles);
textareaRef?.focus();
inputRef?.focus();
}
function handlePromptLoadComplete(placeholderId: string, result: GetPromptResult) {
@@ -426,39 +428,30 @@
onUploadedFilesChange?.(uploadedFiles);
}
function handlePromptPickerClose() {
isPromptPickerOpen = false;
promptSearchQuery = '';
textareaRef?.focus();
// Deferred so the closing popover's focus scope tears down first -
// bits-ui yanks a synchronous focus() back into the still-mounted popover.
function refocusInput() {
queueMicrotask(() => inputRef?.focus());
}
function handleInlineResourcePickerClose() {
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
textareaRef?.focus();
}
// Splice the mention link in place of the `@<query>` token. Uses the
// live cursor, not a stale snapshot - the token may have been edited.
function handleMentionSelect(entry: FileMentionEntry) {
const cursor = inputRef?.getCaretOffset() ?? value.length;
const token = findMentionToken(value, cursor);
if (!token) return;
function handleInlineResourceSelect() {
if (value.startsWith(RESOURCE_TRIGGER_PREFIX)) {
value = '';
onValueChange?.('');
}
const built = buildMentionInsertion(entry, value, token);
if (!built) return;
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
textareaRef?.focus();
}
value = built.newValue;
onValueChange?.(built.newValue);
function handleBrowseResources() {
isInlineResourcePickerOpen = false;
resourceSearchQuery = '';
if (value.startsWith(RESOURCE_TRIGGER_PREFIX)) {
value = '';
onValueChange?.('');
}
isResourceDialogOpen = true;
// bind:value applies on the next microtask; restore the caret after.
queueMicrotask(() => {
inputRef?.focus();
inputRef?.setCaretOffset(built.caretOffset);
});
}
async function handleMicClick() {
@@ -503,19 +496,32 @@
>
<ChatFormPickers
bind:this={pickersRef}
{isPromptPickerOpen}
{promptSearchQuery}
{isInlineResourcePickerOpen}
{resourceSearchQuery}
onPromptPickerClose={handlePromptPickerClose}
onInlineResourcePickerClose={handleInlineResourcePickerClose}
onInlineResourceSelect={handleInlineResourceSelect}
isCommandPickerOpen={pickers.isCommandPickerOpen}
commandQuery={pickers.commandQuery}
commands={pickers.availableCommands}
onCommandPickerClose={pickers.handleCommandPickerClose}
onCommandSelect={pickers.handleCommandSelect}
isPromptPickerOpen={pickers.isPromptPickerOpen}
promptSearchQuery={pickers.promptSearchQuery}
isMentionPickerOpen={pickers.isMentionPickerOpen}
mentionQuery={pickers.mentionQuery}
{mentionAnchor}
scopePath={pickers.mentionScopePath}
onPromptPickerClose={pickers.handlePromptPickerClose}
onMentionPickerClose={pickers.handleMentionPickerClose}
onMentionOpened={() => inputRef?.focus()}
onMentionSelect={handleMentionSelect}
onPromptLoadStart={handlePromptLoadStart}
onPromptLoadComplete={handlePromptLoadComplete}
onPromptLoadError={handlePromptLoadError}
onInlineResourceBrowse={handleBrowseResources}
/>
<div
bind:this={mentionAnchor}
class="pointer-events-none absolute top-0 right-0 left-0 h-px"
aria-hidden="true"
></div>
<div
class="{INPUT_CLASSES} overflow-hidden rounded-4xl md:rounded-3xl backdrop-blur-md {disabled
? 'cursor-not-allowed opacity-60'
@@ -534,17 +540,17 @@
<div
class="flex-column relative min-h-12 items-center rounded-4xl md:rounded-3xl py-2 pb-2.25 shadow-sm transition-all focus-within:shadow-md md:py-3!"
onpaste={handlePaste}
>
<ChatFormTextarea
class="px-5 py-1.5 md:pt-0"
bind:this={textareaRef}
bind:this={inputRef}
bind:value
onKeydown={handleKeydown}
onInput={() => {
handleInput();
pickers.handleInput();
onValueChange?.(value);
}}
onPaste={handlePaste}
{disabled}
{placeholder}
/>
@@ -574,7 +580,7 @@
onMicClick={handleMicClick}
{onStop}
onSystemPromptClick={() => onSystemPromptClick?.({ message: value, files: uploadedFiles })}
onMcpPromptClick={showMcpPromptButton ? () => (isPromptPickerOpen = true) : undefined}
onMcpPromptClick={showMcpPromptButton ? () => pickers.openPromptPicker() : undefined}
onMcpResourcesClick={() => (isResourceDialogOpen = true)}
/>
</div>
@@ -585,8 +591,12 @@
{#if toolsStore.builtinTools.length > 0}
<ChatFormWorkingDirectory
directory={cwd}
isOpen={pickers.isWorkingDirectoryPickerOpen}
bind:query={pickers.workingDirectoryQuery}
customAnchor={mentionAnchor}
onChange={handleWorkingDirectoryChange}
onClose={refocusInput}
onClose={pickers.handleWorkingDirectoryClose}
onOpen={pickers.handleWorkingDirectoryOpen}
{disabled}
/>
{/if}
@@ -0,0 +1,135 @@
<script lang="ts">
import { FolderOpen, Sparkles } from '@lucide/svelte';
import { MODEL_SELECTOR_ICON } from '$lib/constants';
import { usePickerNavigation } from '$lib/hooks/use-picker-navigation.svelte';
import { ChatFormCommandAction } from '$lib/enums';
import type { ChatFormCommand } from '$lib/types';
import {
ChatFormPickerList,
ChatFormPickerListItem,
ChatFormPickerPopover
} from '$lib/components/app/chat';
/**
* Slash-command picker; `query` (typed after `/`) filters the commands.
* The parent owns the "dismissed token, don't act until it changes"
* snapshot, so this picker just renders and reports selection.
*/
interface Props {
class?: string;
isOpen: boolean;
query: string;
commands: ChatFormCommand[];
onClose: () => void;
onSelect: (command: ChatFormCommand) => void;
}
let { class: className = '', isOpen, query, commands, onClose, onSelect }: Props = $props();
const commandIcon: Record<ChatFormCommandAction, typeof Sparkles> = {
[ChatFormCommandAction.PROMPT]: Sparkles,
[ChatFormCommandAction.CWD]: FolderOpen,
[ChatFormCommandAction.MODEL]: MODEL_SELECTOR_ICON
};
const trimmedQuery = $derived((query ?? '').trim().toLowerCase());
const filteredCommands = $derived(
trimmedQuery
? commands.filter(
(c) =>
c.name.toLowerCase().includes(trimmedQuery) ||
c.description.toLowerCase().includes(trimmedQuery) ||
(c.keywords ?? []).some((k) => k.toLowerCase().includes(trimmedQuery))
)
: commands
);
function firstEnabledIndex(): number {
return filteredCommands.findIndex((c) => !c.disabled);
}
function stepEnabled(from: number, dir: number): number {
const n = filteredCommands.length;
if (n === 0) return -1;
for (let i = 1; i <= n; i++) {
const idx = (from + dir * i + n) % n;
if (!filteredCommands[idx].disabled) return idx;
}
return -1;
}
const nav = usePickerNavigation({
isOpen: () => isOpen,
count: () => filteredCommands.length,
step: (from, dir) => (from < 0 ? firstEnabledIndex() : stepEnabled(from, dir)),
onClose: () => onClose(),
onSelect: (index) => handleSelect(filteredCommands[index])
});
$effect(() => {
if (isOpen) {
nav.reset(firstEnabledIndex());
}
});
$effect(() => {
if (nav.hoveredIndex < 0 || nav.hoveredIndex >= filteredCommands.length) {
nav.reset(firstEnabledIndex());
return;
}
if (filteredCommands[nav.hoveredIndex].disabled) {
nav.reset(firstEnabledIndex());
}
});
function handleSelect(command: ChatFormCommand) {
if (command.disabled) return;
onSelect(command);
onClose();
}
export function handleKeydown(event: KeyboardEvent): boolean {
return nav.handleKeydown(event);
}
</script>
<ChatFormPickerPopover
bind:isOpen
class={className}
srLabel="Open command picker"
{onClose}
onKeydown={handleKeydown}
>
<ChatFormPickerList
items={filteredCommands}
isLoading={false}
selectedIndex={nav.hoveredIndex}
showSearchInput={false}
searchQuery={query ?? ''}
emptyMessage="No matching command"
itemKey={(command) => command.name}
scrollTrigger={nav.scrollTrigger}
>
{#snippet item(command, index, isSelected)}
{@const Icon = commandIcon[command.action]}
<ChatFormPickerListItem
dataIndex={index}
{isSelected}
disabled={command.disabled}
onclick={() => handleSelect(command)}
onmouseenter={() => {
if (!command.disabled) nav.setHover(index);
}}
>
<Icon class="mt-0.5 h-4 w-4 shrink-0 text-muted-foreground" />
<div class="flex min-w-0 flex-1 flex-col">
<span class="font-mono text-sm font-medium">/{command.name}</span>
<span class="min-w-0 flex-1 truncate text-left text-xs text-muted-foreground">
{command.description}
</span>
</div>
</ChatFormPickerListItem>
{/snippet}
</ChatFormPickerList>
</ChatFormPickerPopover>
@@ -0,0 +1,258 @@
<script lang="ts">
import { File, Folder } from '@lucide/svelte';
import { abbreviateHome, runGlobSearchWithChildren, type GlobEntryResult } from '$lib/utils';
import { toolsStore } from '$lib/stores/tools.svelte';
import { BuiltInTool, FileMentionEntryType, GlobSearchType } from '$lib/enums';
import { isMobile } from '$lib/stores/viewport.svelte';
import { config } from '$lib/stores/settings.svelte';
import * as Popover from '$lib/components/ui/popover';
import * as Tooltip from '$lib/components/ui/tooltip';
import HighlightedMatch from '$lib/components/app/forms/HighlightedMatch.svelte';
import { ChatFormPickerList, ChatFormPickerListItem } from '$lib/components/app/chat';
import { useDebouncedSearch } from '$lib/hooks/use-debounced-search.svelte';
import { usePickerNavigation } from '$lib/hooks/use-picker-navigation.svelte';
import type { FileMentionEntry } from '$lib/types';
import {
FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH,
HOME_TILDE,
SEARCH_DEBOUNCE_MS
} from '$lib/constants';
/**
* Floating file/folder mention picker. The chat input is the search
* surface: `query` (typed after `@`) drives a `file_glob_search` tool
* call scoped to `scopePath`. The parent owns the "dismissed token,
* don't re-open until it changes" snapshot.
*/
interface Props {
class?: string;
isOpen: boolean;
query: string;
customAnchor?: HTMLElement | null;
scopePath?: string | null;
onClose: () => void;
onSelect: (entry: FileMentionEntry) => void;
/** Fired when `isOpen` becomes true, so the host can keep focus on the chat input. */
onOpened?: () => void;
}
let {
class: className = '',
isOpen,
query,
customAnchor = null,
scopePath = null,
onClose,
onSelect,
onOpened
}: Props = $props();
const nav = usePickerNavigation({
isOpen: () => isOpen,
count: () => displayedItems.length,
onClose: () => onClose(),
onSelect: (index) => handleSelect(displayedItems[index])
});
// When the server does not expose file_glob_search (started without
// --tools) or the user disabled it, the picker still opens but explains
// why instead of firing searches that would only fail.
const fileSearchKey = $derived(toolsStore.getPermissionKey(BuiltInTool.FILE_GLOB_SEARCH));
const fileSearchEnabled = $derived(
fileSearchKey !== null && toolsStore.isToolEnabled(fileSearchKey)
);
let searchResults = $state<FileMentionEntry[]>([]);
let searchError = $state<string | null>(null);
// Coerce the depth setting to a positive integer; an invalid value
// would otherwise reach the server as max_depth 0 = unlimited.
const searchDepth = $derived.by(() => {
const n = Number(config().mentionSearchMaxDepth);
return Number.isInteger(n) && n > 0 ? n : FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH;
});
const home = $derived(toolsStore.serverHome);
// A smaller window than the WD picker suffices: entries are ranked client-side.
const MENTION_SEARCH_LIMIT = 50;
const search = useDebouncedSearch({
debounceMs: SEARCH_DEBOUNCE_MS,
canRun: () => isOpen && fileSearchEnabled,
getQuery: () => trimmedQuery,
run: async (query, signal, isCurrent) => {
try {
// A trailing path separator targets a directory, so also list its
// children. Accept both `/` and `\`.
const res = await runGlobSearchWithChildren(
query,
scopePath ?? home ?? HOME_TILDE,
searchDepth,
MENTION_SEARCH_LIMIT,
signal,
{ type: GlobSearchType.ALL, descendOnTrailingSeparator: true }
);
if (!isCurrent()) return;
if (res.error) {
searchResults = [];
searchError = res.error;
return;
}
const toEntry = (e: GlobEntryResult): FileMentionEntry => ({
path: e.path,
name: e.name,
type: e.type === 'dir' ? FileMentionEntryType.DIRECTORY : FileMentionEntryType.FILE
});
searchResults = res.entries.map(toEntry);
searchError = null;
} catch (err) {
if (!isCurrent() || signal.aborted) return;
searchResults = [];
searchError = err instanceof Error ? err.message : String(err);
}
}
});
const trimmedQuery = $derived((query ?? '').trim());
const displayedItems = $derived(searchResults);
const emptyMessage = $derived.by(() => {
if (fileSearchKey === null) {
return 'File search is unavailable on this server (started without --tools)';
}
if (!fileSearchEnabled) {
return 'File search is disabled - enable "Search files" in Settings > Tools to use @-mentions';
}
return searchError ? `Search failed - ${searchError}` : 'No matching files or folders';
});
const showTooltip = $derived(!isMobile.current);
$effect(() => {
if (typeof window === 'undefined') return;
void toolsStore.resolveServerHome();
});
$effect(() => {
if (isOpen) {
nav.reset(0);
}
});
$effect(() => {
if (isOpen) onOpened?.();
});
$effect(() => {
const q = (query ?? '').trim();
if (!isOpen || !q || !fileSearchEnabled) {
search.cancel();
searchResults = [];
searchError = null;
return;
}
search.setLoading(true);
search.run(q);
});
function handleSelect(entry: FileMentionEntry) {
onSelect(entry);
onClose();
}
export function handleKeydown(event: KeyboardEvent): boolean {
return nav.handleKeydown(event);
}
</script>
<Popover.Root
open={isOpen}
onOpenChange={(open) => {
if (!open) onClose();
}}
>
<!-- Invisible form-wide trigger: stops bits-ui's outside-click detector
from closing the picker when the user clicks inside the textarea.
We open programmatically via `open={isOpen}`, so it is inert
(tabindex=-1 + pointer-events-none + opacity-0 + aria-hidden).
Positioning comes from `customAnchor` at the form's top edge. -->
<Popover.Trigger
class="pointer-events-none absolute inset-0 opacity-0"
tabindex={-1}
aria-hidden="true"
>
<span class="sr-only">Open file mention picker</span>
</Popover.Trigger>
<Popover.Content
align="start"
side="top"
sideOffset={12}
{customAnchor}
preventScroll={false}
onkeydown={handleKeydown}
onOpenAutoFocus={(event) => event.preventDefault()}
onCloseAutoFocus={(event) => event.preventDefault()}
class={[
'w-[var(--bits-popover-anchor-width)] max-w-none rounded-xl border-border/50 p-0 shadow-xl',
className
]}
>
<ChatFormPickerList
items={displayedItems}
isLoading={search.isSearching}
selectedIndex={nav.hoveredIndex}
showSearchInput={false}
searchQuery={query ?? ''}
{emptyMessage}
itemKey={(entry) => entry.type + ':' + entry.path}
scrollTrigger={nav.scrollTrigger}
>
{#snippet item(entry, index, isSelected)}
<ChatFormPickerListItem
dataIndex={index}
{isSelected}
onclick={() => handleSelect(entry)}
onmouseenter={() => nav.setHover(index)}
>
{@const Icon = entry.type === FileMentionEntryType.DIRECTORY ? Folder : File}
<Icon
class={[
'mt-0.5 h-4 w-4 shrink-0',
entry.type === FileMentionEntryType.DIRECTORY
? 'text-amber-500'
: 'text-muted-foreground'
]}
/>
<div class="flex min-w-0 flex-1 flex-col">
<div class="flex min-w-0 items-center gap-2">
{#if showTooltip}
<Tooltip.Root>
<Tooltip.Trigger>
{#snippet child({ props })}
<span {...props} class="truncate text-sm font-medium">{entry.name}</span>
{/snippet}
</Tooltip.Trigger>
<Tooltip.Content>
<p>{entry.path}</p>
</Tooltip.Content>
</Tooltip.Root>
{:else}
<span class="truncate text-sm font-medium">{entry.name}</span>
{/if}
<span
class="shrink-0 rounded-full bg-muted px-1.5 py-0.5 font-mono text-[9px] uppercase tracking-wide text-muted-foreground"
>
{entry.type}
</span>
</div>
<span class="min-w-0 flex-1 truncate font-mono text-left text-xs">
<HighlightedMatch text={abbreviateHome(entry.path, home)} query={trimmedQuery} />
</span>
</div>
</ChatFormPickerListItem>
{/snippet}
</ChatFormPickerList>
</Popover.Content>
</Popover.Root>
@@ -2,6 +2,7 @@
import type { Snippet } from 'svelte';
import { SearchInput } from '$lib/components/app';
import ScrollArea from '$lib/components/ui/scroll-area/scroll-area.svelte';
import { useScrollActiveRow } from '$lib/hooks/use-scroll-active-row.svelte';
import { CHAT_FORM_POPOVER_MAX_HEIGHT } from '$lib/constants';
interface Props {
@@ -11,11 +12,19 @@
searchQuery: string;
showSearchInput: boolean;
searchPlaceholder?: string;
// Omit to distinguish "haven't searched yet" from "search returned nothing".
emptyMessage?: string;
autofocus?: boolean;
inputRef?: HTMLInputElement | null;
onSearchClose?: () => void;
itemKey: (item: T, index: number) => string;
item: Snippet<[T, number, boolean]>;
skeleton?: Snippet;
skeletonCount?: number;
footer?: Snippet;
// Counter bumped by the picker on keyboard nav; scrolls the selected
// row into view without scrolling on hover or result replacement.
scrollTrigger?: number;
}
let {
@@ -25,49 +34,69 @@
searchQuery = $bindable(),
showSearchInput,
searchPlaceholder = 'Search...',
emptyMessage = 'No items available',
emptyMessage,
autofocus = false,
inputRef = $bindable(null),
onSearchClose,
itemKey,
item,
skeleton,
footer
skeletonCount = 6,
footer,
scrollTrigger
}: Props = $props();
let listContainer = $state<HTMLDivElement | null>(null);
$effect(() => {
if (listContainer && selectedIndex >= 0 && selectedIndex < items.length) {
const selectedElement = listContainer.querySelector(
`[data-picker-index="${selectedIndex}"]`
) as HTMLElement;
let listPaddingTop = $derived(
showSearchInput ? (isLoading || items.length > 0 ? 'pt-13' : 'pt-10') : ''
);
if (selectedElement) {
selectedElement.scrollIntoView({
behavior: 'smooth',
block: 'center',
inline: 'nearest'
});
}
}
// selectedIndex/items.length are untracked so hover and result replacement
// never re-fire the scroll; keyboard nav is the only path that bumps the trigger.
useScrollActiveRow({
getTrigger: () => scrollTrigger,
getContainer: () => listContainer,
getIndex: () => selectedIndex,
getCount: () => items.length,
dataIndex: 'picker'
});
</script>
<ScrollArea>
{#if showSearchInput}
<div class="absolute top-0 right-0 left-0 z-10 p-2 pb-0">
<SearchInput placeholder={searchPlaceholder} bind:value={searchQuery} />
<SearchInput
{autofocus}
placeholder={searchPlaceholder}
bind:value={searchQuery}
bind:ref={inputRef}
onClose={onSearchClose}
/>
</div>
{/if}
<div
bind:this={listContainer}
class={[`${CHAT_FORM_POPOVER_MAX_HEIGHT} p-2`, showSearchInput && 'pt-13']}
>
<div bind:this={listContainer} class={[`${CHAT_FORM_POPOVER_MAX_HEIGHT} p-2`, listPaddingTop]}>
{#if isLoading}
{#if skeleton}
{@render skeleton()}
{:else}
<div aria-busy="true" aria-live="polite" class="flex flex-col">
{#each { length: skeletonCount } as _, rowIndex (rowIndex)}
<div class="flex items-start gap-3 rounded-lg px-3 py-2">
<div class="mt-0.5 size-4 shrink-0 animate-pulse rounded-md bg-muted/60"></div>
<div class="flex min-w-0 flex-1 flex-col">
<div class="h-5 w-2/5 animate-pulse rounded-sm bg-muted/60"></div>
<div class="h-4 w-1/3 animate-pulse rounded-sm bg-muted/40"></div>
</div>
</div>
{/each}
</div>
{/if}
{:else if items && items.length === 0}
{#if emptyMessage}
<div class="py-6 text-center text-sm text-muted-foreground">{emptyMessage}</div>
{/if}
{:else if items.length === 0}
<div class="py-6 text-center text-sm text-muted-foreground">{emptyMessage}</div>
{:else}
{#each items as itemData, index (itemKey(itemData, index))}
{@render item(itemData, index, index === selectedIndex)}
@@ -3,21 +3,34 @@
interface Props {
isSelected?: boolean;
disabled?: boolean;
onclick: () => void;
onmouseenter?: () => void;
dataIndex?: number;
children: Snippet;
class?: string;
}
let { isSelected = false, onclick, dataIndex, children }: Props = $props();
let {
class: className = '',
isSelected = false,
disabled = false,
onclick,
onmouseenter,
dataIndex,
children
}: Props = $props();
</script>
<button
type="button"
data-picker-index={dataIndex}
{disabled}
{onclick}
{onmouseenter}
class="flex w-full cursor-pointer items-start gap-3 rounded-lg px-3 py-2 text-left hover:bg-accent/50 {isSelected
? 'bg-accent/50'
: ''}"
: ''} {disabled ? 'cursor-not-allowed opacity-50' : ''} {className}"
>
{@render children()}
</button>
@@ -42,6 +42,7 @@
align="start"
sideOffset={12}
class="w-[var(--bits-popover-anchor-width)] max-w-none rounded-xl border-border/50 p-0 shadow-xl {className}"
preventScroll={false}
onkeydown={onKeydown}
onOpenAutoFocus={(event) => event.preventDefault()}
>
@@ -45,6 +45,9 @@
let promptArgs = $state<Record<string, string>>({});
let selectedIndex = $state(0);
let internalSearchQuery = $state('');
// Bumped on ArrowUp/ArrowDown only, so the list scrolls on keyboard
// nav but not on hover or result changes.
let scrollTrigger = $state(0);
let promptError = $state<string | null>(null);
let selectedIndexBeforeArgumentForm = $state<number | null>(null);
@@ -295,6 +298,7 @@
event.preventDefault();
if (filteredPrompts.length > 0) {
selectedIndex = (selectedIndex + 1) % filteredPrompts.length;
scrollTrigger++;
}
return true;
@@ -304,6 +308,7 @@
event.preventDefault();
if (filteredPrompts.length > 0) {
selectedIndex = selectedIndex === 0 ? filteredPrompts.length - 1 : selectedIndex - 1;
scrollTrigger++;
}
return true;
@@ -400,6 +405,7 @@
searchPlaceholder="Search prompts..."
emptyMessage="No MCP prompts available"
itemKey={(prompt) => prompt.serverName + ':' + prompt.name}
{scrollTrigger}
>
{#snippet item(prompt, index, isSelected)}
{@const server = serverSettingsMap.get(prompt.serverName)}
@@ -1,237 +0,0 @@
<script lang="ts">
import { conversationsStore } from '$lib/stores/conversations.svelte';
import { mcpStore } from '$lib/stores/mcp.svelte';
import { mcpResourceStore } from '$lib/stores/mcp-resources.svelte';
import { KeyboardKey } from '$lib/enums';
import type { MCPResourceInfo, MCPServerSettingsEntry } from '$lib/types';
import { SvelteMap } from 'svelte/reactivity';
import { FolderOpen } from '@lucide/svelte';
import { Button } from '$lib/components/ui/button';
import {
ChatFormPickerPopover,
ChatFormPickerList,
ChatFormPickerListItem,
ChatFormPickerItemHeader,
ChatFormPickerListItemSkeleton
} from '$lib/components/app/chat';
interface Props {
class?: string;
isOpen?: boolean;
searchQuery?: string;
onClose?: () => void;
onResourceSelect?: (resource: MCPResourceInfo) => void;
onBrowse?: () => void;
}
let {
class: className = '',
isOpen = false,
searchQuery = '',
onClose,
onResourceSelect,
onBrowse
}: Props = $props();
let resources = $state<MCPResourceInfo[]>([]);
let isLoading = $state(false);
let selectedIndex = $state(0);
let internalSearchQuery = $state('');
let serverSettingsMap = $derived.by(() => {
const servers = mcpStore.getServers();
const map = new SvelteMap<string, MCPServerSettingsEntry>();
for (const server of servers) {
map.set(server.id, server);
}
return map;
});
$effect(() => {
if (isOpen) {
loadResources();
selectedIndex = 0;
}
});
$effect(() => {
if (filteredResources.length > 0 && selectedIndex >= filteredResources.length) {
selectedIndex = 0;
}
});
async function loadResources() {
isLoading = true;
try {
const perChatOverrides = conversationsStore.getAllMcpServerOverrides();
const initialized = await mcpStore.ensureInitialized(perChatOverrides);
if (!initialized) {
resources = [];
return;
}
await mcpStore.fetchAllResources();
resources = mcpResourceStore.getAllResourceInfos();
} catch (error) {
console.error('[ChatFormPickerMcpResources] Failed to load resources:', error);
resources = [];
} finally {
isLoading = false;
}
}
function handleResourceClick(resource: MCPResourceInfo) {
mcpStore.attachResource(resource.uri);
onResourceSelect?.(resource);
onClose?.();
}
function isResourceAttached(uri: string): boolean {
return mcpResourceStore.isAttached(uri);
}
export function handleKeydown(event: KeyboardEvent): boolean {
if (!isOpen) return false;
if (event.key === KeyboardKey.ESCAPE) {
event.preventDefault();
onClose?.();
return true;
}
if (event.key === KeyboardKey.ARROW_DOWN) {
event.preventDefault();
if (filteredResources.length > 0) {
selectedIndex = (selectedIndex + 1) % filteredResources.length;
}
return true;
}
if (event.key === KeyboardKey.ARROW_UP) {
event.preventDefault();
if (filteredResources.length > 0) {
selectedIndex = selectedIndex === 0 ? filteredResources.length - 1 : selectedIndex - 1;
}
return true;
}
if (event.key === KeyboardKey.ENTER) {
event.preventDefault();
if (filteredResources[selectedIndex]) {
handleResourceClick(filteredResources[selectedIndex]);
}
return true;
}
return false;
}
let filteredResources = $derived.by(() => {
const sortedServers = mcpStore.getServers();
const serverOrderMap = new Map(sortedServers.map((server, index) => [server.id, index]));
const sortedResources = [...resources].sort((a, b) => {
const orderA = serverOrderMap.get(a.serverName) ?? Number.MAX_SAFE_INTEGER;
const orderB = serverOrderMap.get(b.serverName) ?? Number.MAX_SAFE_INTEGER;
return orderA - orderB;
});
const query = (searchQuery || internalSearchQuery).toLowerCase();
if (!query) return sortedResources;
return sortedResources.filter(
(resource) =>
resource.name.toLowerCase().includes(query) ||
resource.title?.toLowerCase().includes(query) ||
resource.description?.toLowerCase().includes(query) ||
resource.uri.toLowerCase().includes(query)
);
});
let showSearchInput = $derived(resources.length > 3);
</script>
<ChatFormPickerPopover
bind:isOpen
class={className}
srLabel="Open resource picker"
{onClose}
onKeydown={handleKeydown}
>
<ChatFormPickerList
items={filteredResources}
{isLoading}
{selectedIndex}
bind:searchQuery={internalSearchQuery}
{showSearchInput}
searchPlaceholder="Search resources..."
emptyMessage="No MCP resources available"
itemKey={(resource) => resource.serverName + ':' + resource.uri}
>
{#snippet item(resource, index, isSelected)}
{@const server = serverSettingsMap.get(resource.serverName)}
{@const serverLabel = server ? mcpStore.getServerLabel(server) : resource.serverName}
<ChatFormPickerListItem
dataIndex={index}
{isSelected}
onclick={() => handleResourceClick(resource)}
>
<ChatFormPickerItemHeader
{server}
{serverLabel}
title={resource.title || resource.name}
description={resource.description}
>
{#snippet titleExtra()}
{#if isResourceAttached(resource.uri)}
<span
class="inline-flex items-center rounded-full bg-primary/10 px-1.5 py-0.5 text-[10px] font-medium text-primary"
>
attached
</span>
{/if}
{/snippet}
{#snippet subtitle()}
<p class="mt-0.5 truncate text-xs text-muted-foreground/60">
{resource.uri}
</p>
{/snippet}
</ChatFormPickerItemHeader>
</ChatFormPickerListItem>
{/snippet}
{#snippet skeleton()}
<ChatFormPickerListItemSkeleton />
{/snippet}
{#snippet footer()}
{#if onBrowse && resources.length > 3}
<Button
class="fixed right-3 bottom-3"
type="button"
onclick={onBrowse}
variant="secondary"
size="sm"
>
<FolderOpen class="h-3 w-3" />
Browse all
</Button>
{/if}
{/snippet}
</ChatFormPickerList>
</ChatFormPickerPopover>
@@ -1,16 +1,30 @@
<script lang="ts">
import ChatFormCommandPicker from './ChatFormCommandPicker.svelte';
import ChatFormMentionPicker from './ChatFormMentionPicker.svelte';
import ChatFormPickerMcpPrompts from './ChatFormPickerMcpPrompts/ChatFormPickerMcpPrompts.svelte';
import ChatFormPickerMcpResources from './ChatFormPickerMcpResources.svelte';
import type { GetPromptResult, MCPPromptInfo } from '$lib/types';
import type {
ChatFormCommand,
FileMentionEntry,
GetPromptResult,
MCPPromptInfo
} from '$lib/types';
interface Props {
isCommandPickerOpen?: boolean;
commandQuery?: string;
commands?: ChatFormCommand[];
isPromptPickerOpen?: boolean;
promptSearchQuery?: string;
isInlineResourcePickerOpen?: boolean;
resourceSearchQuery?: string;
isMentionPickerOpen?: boolean;
mentionQuery?: string;
mentionAnchor?: HTMLElement | null;
scopePath?: string | null;
onCommandPickerClose?: () => void;
onCommandSelect?: (command: ChatFormCommand) => void;
onPromptPickerClose?: () => void;
onInlineResourcePickerClose?: () => void;
onInlineResourceSelect?: () => void;
onMentionPickerClose?: () => void;
onMentionOpened?: () => void;
onMentionSelect?: (entry: FileMentionEntry) => void;
onPromptLoadStart?: (
placeholderId: string,
promptInfo: MCPPromptInfo,
@@ -18,36 +32,44 @@
) => void;
onPromptLoadComplete?: (placeholderId: string, result: GetPromptResult) => void;
onPromptLoadError?: (placeholderId: string, error: string) => void;
onInlineResourceBrowse?: () => void;
}
let {
isCommandPickerOpen,
commandQuery,
commands = [],
onCommandPickerClose,
onCommandSelect,
isPromptPickerOpen,
promptSearchQuery,
isInlineResourcePickerOpen,
resourceSearchQuery,
isMentionPickerOpen,
mentionQuery,
mentionAnchor,
scopePath,
onPromptPickerClose,
onInlineResourcePickerClose,
onInlineResourceSelect,
onMentionPickerClose,
onMentionOpened,
onMentionSelect,
onPromptLoadStart,
onPromptLoadComplete,
onPromptLoadError,
onInlineResourceBrowse
onPromptLoadError
}: Props = $props();
let commandPickerRef: ChatFormCommandPicker | undefined = $state(undefined);
let promptPickerRef: ChatFormPickerMcpPrompts | undefined = $state(undefined);
let resourcePickerRef: ChatFormPickerMcpResources | undefined = $state(undefined);
let mentionPickerRef: ChatFormMentionPicker | undefined = $state(undefined);
/**
* Delegates keyboard events to the active picker child.
* Returns true if the event was handled.
*/
/** Delegate keyboard events to the active picker child; true if handled. */
export function handleKeydown(event: KeyboardEvent): boolean {
if (isCommandPickerOpen && commandPickerRef?.handleKeydown(event)) {
return true;
}
if (isPromptPickerOpen && promptPickerRef?.handleKeydown(event)) {
return true;
}
if (isInlineResourcePickerOpen && resourcePickerRef?.handleKeydown(event)) {
if (isMentionPickerOpen && mentionPickerRef?.handleKeydown(event)) {
return true;
}
@@ -55,6 +77,15 @@
}
</script>
<ChatFormCommandPicker
bind:this={commandPickerRef}
isOpen={isCommandPickerOpen ?? false}
query={commandQuery ?? ''}
{commands}
onClose={onCommandPickerClose ?? (() => {})}
onSelect={onCommandSelect ?? (() => {})}
/>
<ChatFormPickerMcpPrompts
bind:this={promptPickerRef}
isOpen={isPromptPickerOpen}
@@ -65,11 +96,13 @@
{onPromptLoadError}
/>
<ChatFormPickerMcpResources
bind:this={resourcePickerRef}
isOpen={isInlineResourcePickerOpen}
searchQuery={resourceSearchQuery}
onClose={onInlineResourcePickerClose}
onResourceSelect={onInlineResourceSelect}
onBrowse={onInlineResourceBrowse}
<ChatFormMentionPicker
bind:this={mentionPickerRef}
isOpen={isMentionPickerOpen ?? false}
query={mentionQuery ?? ''}
customAnchor={mentionAnchor}
scopePath={scopePath ?? null}
onClose={onMentionPickerClose ?? (() => {})}
onOpened={onMentionOpened}
onSelect={onMentionSelect ?? (() => {})}
/>
@@ -28,11 +28,10 @@
onMount(() => {
if (textareaElement) {
autoResizeTextarea(textareaElement);
textareaElement.focus();
textareaElement.focus({ preventScroll: true });
}
});
// Expose the textarea element for external access
export function getElement() {
return textareaElement;
}
@@ -48,6 +47,16 @@
textareaElement.style.height = '1rem';
}
}
// Plain-text caret offsets for the picker/paste/mention-splice flows.
export function getCaretOffset(): number {
if (!textareaElement) return 0;
return textareaElement.selectionStart ?? textareaElement.value.length;
}
export function setCaretOffset(offset: number) {
textareaElement?.setSelectionRange(offset, offset);
}
</script>
<div class="flex-1 {className}">
@@ -1,7 +1,5 @@
<script lang="ts">
import { FolderOpen } from '@lucide/svelte';
import { untrack } from 'svelte';
import { SvelteMap } from 'svelte/reactivity';
import { ToolsService } from '$lib/services/tools.service';
import { toolsStore } from '$lib/stores/tools.svelte';
import { BuiltInTool, GlobSearchType, KeyboardKey } from '$lib/enums';
@@ -10,23 +8,22 @@
buildCaseInsensitiveGlob,
joinPath,
lastPathSegment,
rankEntries,
splitPathQuery,
runGlobSearchWithChildren,
type GlobEntry
} from '$lib/utils';
import { debounce } from '$lib/utils/debounce';
import * as Popover from '$lib/components/ui/popover';
import SearchInput from '$lib/components/app/forms/SearchInput.svelte';
import { useDebouncedSearch } from '$lib/hooks/use-debounced-search.svelte';
import { usePickerNavigation } from '$lib/hooks/use-picker-navigation.svelte';
import { useScrollActiveRow } from '$lib/hooks/use-scroll-active-row.svelte';
import ChatFormWorkingDirectoryChip from './ChatFormWorkingDirectoryChip.svelte';
import ChatFormWorkingDirectoryResultsList from './ChatFormWorkingDirectoryResultsList.svelte';
import {
DEFAULT_MOBILE_BREAKPOINT,
GLOB_WILDCARD,
HOME_TILDE,
MAX_RESULTS_SHOWN,
NATIVE_LIMIT,
NATIVE_MAX_DEPTH,
PATH_NAV_MAX_DEPTH,
SEARCH_DEBOUNCE_MS,
SEARCH_LIMIT,
SEARCH_MAX_DEPTH
@@ -39,228 +36,147 @@
class?: string;
disabled?: boolean;
directory?: string | null;
/** Controlled open state; the host owns it so the chip click and the
* `/cwd` slash command open the picker through the same path. */
isOpen: boolean;
/** Two-way bound query, kept in sync with the text after `/cwd `. */
query: string;
/** Anchor at the form's top edge so the popover floats above the box. */
customAnchor?: HTMLElement | null;
onChange?: (directory: string | null) => void;
/**
* Lets the host refocus the chat input so typing can resume without
* an extra click after the popover closes.
*/
/** Lets the host refocus the chat input after the popover closes. */
onClose?: () => void;
/** Fired when the chip is clicked so the host can open the picker. */
onOpen?: () => void;
}
let {
class: className = '',
disabled = false,
directory = $bindable(null),
directory = null,
isOpen,
query = $bindable(''),
customAnchor = null,
onChange,
onClose
onClose,
onOpen
}: Props = $props();
// File System Access API is opt-in: when available (Chrome / Edge / Opera) the popover
// exposes a "Browse" button that opens the native folder picker. When unavailable the
// popover still works via the text input - no alerts, no upload semantics.
// File System Access API is opt-in (Chrome / Edge / Opera): the popover
// exposes a "Browse" button only when available.
const pickerSupported =
typeof window !== 'undefined' && typeof window.showDirectoryPicker === 'function';
// Popover open state; the element handles outside-click and Escape.
let isOpen = $state(false);
let inputValue = $state('');
let searchInputRef: HTMLInputElement | null = $state(null);
let queryResults = $state<string[]>([]);
let isSearching = $state(false);
let searchError = $state<string | null>(null);
let hoveredIndex = $state(-1);
// Bumped only by ArrowUp/ArrowDown handlers; the list scrolls the
// highlighted row into view only via this trigger, never on hover.
let scrollTrigger = $state(0);
let listContainer = $state<HTMLDivElement | null>(null);
// Absolute home directory on the server, resolved once per session by
// the tools store. Anchors both the search scope and the chip's `~`
// abbreviation.
const nav = usePickerNavigation({
isOpen: () => isOpen,
count: () => queryResults.length,
onClose: closePicker,
onSelect: (index) => commit(queryResults[index])
});
let homeBase = $derived(toolsStore.serverHome);
// AbortController + sequence counter to discard stale responses when the user
// keeps typing; a newer call aborts the previous one. The sequence counter
// also covers the gap between abort and the catch handler.
let searchController: AbortController | null = null;
let searchSeq = 0;
// Cache of the last file_glob_search result per (parent, include, max_depth),
// so repeated queries in the same directory don't re-walk the tree. Entering
// a directory hits it every time: the children listed for an exactly typed
// segment are what the next keystroke, the trailing slash, asks for again.
const SEARCH_CACHE_TTL_MS = 2000;
const searchCache = new SvelteMap<string, { results: GlobEntry[]; base: string; at: number }>();
const runSearch = debounce((query: string) => {
void doSearch(query);
}, SEARCH_DEBOUNCE_MS);
// Resolve home eagerly on mount so the chip can abbreviate before the
// user opens the picker. resolveServerHome() is cached, so repeat calls
// (e.g. from handleOpenChange) are no-ops.
// Resolve home eagerly so the chip can abbreviate before the picker opens.
$effect(() => {
if (typeof window === 'undefined') return;
void toolsStore.resolveServerHome();
});
// Auto-focus the search input when the popover opens.
// HTML `autofocus` is unreliable on dynamically shown elements, so we
// use a microtask (0ms setTimeout) after the effect flushes.
// HTML `autofocus` is unreliable on dynamically shown elements.
$effect(() => {
if (!isOpen) return;
setTimeout(() => searchInputRef?.focus(), FOCUS_DELAY_MS);
});
let lastScrollTrigger: number | null = null;
// hoveredIndex/queryResults are untracked so hover and result replacement
// never re-fire the scroll; keyboard nav is the only path that bumps the trigger
$effect(() => {
if (scrollTrigger === lastScrollTrigger) return;
lastScrollTrigger = scrollTrigger;
untrack(() => {
if (!listContainer) return;
if (hoveredIndex < 0 || hoveredIndex >= queryResults.length) return;
const selectedElement = listContainer.querySelector(
`[data-result-index="${hoveredIndex}"]`
) as HTMLElement | null;
selectedElement?.scrollIntoView({ block: 'nearest', inline: 'nearest' });
});
if (!isOpen) return;
const q = query.trim();
nav.reset(-1);
if (q) {
search.run(q);
} else {
search.cancel();
queryResults = [];
searchError = null;
nav.reset(-1);
searchScope = homeBase ?? HOME_TILDE;
}
});
function cancelSearch() {
searchController?.abort();
searchSeq++;
isSearching = false;
}
useScrollActiveRow({
getTrigger: () => nav.scrollTrigger,
getContainer: () => listContainer,
getIndex: () => nav.hoveredIndex,
getCount: () => queryResults.length,
dataIndex: 'result'
});
// Effective directory the current search runs against (shown in the
// footer); updated by doSearch, including when an exactly-typed
// directory is "entered".
let searchScope = $state(HOME_TILDE);
// Runs a directory listing through the cache, so a repeated query in the
// same directory does not re-walk the tree on the server.
async function searchDirs(
path: string,
include: string,
maxDepth: number,
signal: AbortSignal
): Promise<{ base: string; entries: GlobEntry[]; error?: string }> {
const key = `${path}\u0000${include}\u0000${maxDepth}`;
const cached = searchCache.get(key);
if (cached && Date.now() - cached.at < SEARCH_CACHE_TTL_MS) {
return { base: cached.base, entries: cached.results };
}
const res = await ToolsService.executeToolRaw(
BuiltInTool.FILE_GLOB_SEARCH,
{ path, type: GlobSearchType.DIR, include, max_depth: maxDepth, limit: SEARCH_LIMIT },
signal
);
if (typeof res.error === 'string') return { base: '', entries: [], error: res.error };
const base = typeof res.base === 'string' ? res.base : '';
const entries = Array.isArray(res.entries) ? (res.entries as GlobEntry[]) : [];
const now = Date.now();
for (const [k, v] of searchCache) {
if (now - v.at >= SEARCH_CACHE_TTL_MS) searchCache.delete(k);
}
searchCache.set(key, { results: entries, base, at: now });
return { base, entries };
}
async function doSearch(query: string) {
const trimmed = query.trim();
if (!trimmed) {
queryResults = [];
searchError = null;
isSearching = false;
hoveredIndex = -1;
searchScope = homeBase ?? HOME_TILDE;
return;
}
cancelSearch();
const controller = new AbortController();
searchController = controller;
const mySeq = ++searchSeq;
const pathQuery = splitPathQuery(trimmed);
isSearching = true;
try {
// A generous limit is requested because ranking happens
// client-side; only the top 20 are shown.
const searchPath = pathQuery ? pathQuery.parent : (homeBase ?? HOME_TILDE);
const include = pathQuery
? pathQuery.last
? buildCaseInsensitiveGlob(pathQuery.last)
: GLOB_WILDCARD
: buildCaseInsensitiveGlob(trimmed);
const maxDepth = pathQuery ? PATH_NAV_MAX_DEPTH : SEARCH_MAX_DEPTH;
const res = await searchDirs(searchPath, include, maxDepth, controller.signal);
if (mySeq !== searchSeq) return;
if (res.error) {
// An exactly-typed directory is "entered": the shared search lists its
// children too, so path navigation does not require a trailing slash.
const search = useDebouncedSearch({
debounceMs: SEARCH_DEBOUNCE_MS,
canRun: () => isOpen,
getQuery: () => query.trim(),
run: async (q, signal, isCurrent) => {
const trimmed = q.trim();
if (!trimmed) {
queryResults = [];
hoveredIndex = -1;
searchError = res.error;
searchError = null;
nav.reset(-1);
searchScope = homeBase ?? HOME_TILDE;
return;
}
const { base, entries } = res;
const ranked = rankEntries(entries, pathQuery?.last ?? trimmed);
let results = ranked.map((e) => joinPath(base, e.path));
searchScope = pathQuery ? pathQuery.parent : (homeBase ?? HOME_TILDE);
// An exactly-typed directory is "entered": list its children too,
// so path navigation doesn't require a trailing slash.
const last = pathQuery?.last;
const exact = last
? ranked.find((e) => lastPathSegment(e.path).toLowerCase() === last.toLowerCase())
: undefined;
if (exact) {
const exactDir = joinPath(base, exact.path);
const childRes = await searchDirs(
exactDir,
GLOB_WILDCARD,
PATH_NAV_MAX_DEPTH,
controller.signal
try {
// Generous limit: ranking is client-side, only the top
// MAX_RESULTS_SHOWN are shown.
const res = await runGlobSearchWithChildren(
trimmed,
homeBase ?? HOME_TILDE,
SEARCH_MAX_DEPTH,
SEARCH_LIMIT,
signal,
{ type: GlobSearchType.DIR }
);
if (mySeq !== searchSeq) return;
if (!childRes.error) {
const children = childRes.entries
.map((e) => joinPath(childRes.base, e.path))
.sort((a, b) => a.localeCompare(b));
results = [...results, ...children];
searchScope = exactDir;
if (!isCurrent()) return;
if (res.error) {
queryResults = [];
nav.reset(-1);
searchError = res.error;
return;
}
searchScope = res.exactDir ?? res.args.path;
queryResults = res.entries.map((e) => e.path).slice(0, MAX_RESULTS_SHOWN);
if (queryResults.length > 0) {
nav.reset(0);
nav.bumpScroll(); // scroll the list back to the top (first item is hovered)
} else {
nav.reset(-1);
}
searchError = null;
} catch (err) {
if (!isCurrent() || signal.aborted) return;
queryResults = [];
nav.reset(-1);
searchError = err instanceof Error ? err.message : String(err);
}
queryResults = results.slice(0, MAX_RESULTS_SHOWN);
hoveredIndex = queryResults.length > 0 ? 0 : -1;
// new results: scroll the list back to the top (first item is hovered)
if (hoveredIndex === 0) scrollTrigger++;
searchError = null;
} catch (err) {
if (mySeq !== searchSeq) return;
queryResults = [];
hoveredIndex = -1;
if (controller.signal.aborted) return;
searchError = err instanceof Error ? err.message : String(err);
} finally {
if (mySeq === searchSeq) isSearching = false;
}
}
// Single funnel for every local close so the host refocus fires
// regardless of which commit/dismiss path ended the interaction.
});
// Single funnel for every local close so the host refocus always fires.
function closePicker() {
isOpen = false;
onClose?.();
}
function commit(path: string) {
directory = path;
onChange?.(path);
closePicker();
}
@@ -268,15 +184,12 @@
function setDirectory(value: string) {
const trimmed = value.trim();
if (!trimmed) return;
directory = trimmed;
onChange?.(trimmed);
}
// Resolve a folder name picked via the browser-native picker (which exposes
// only the leaf name) to a server-side absolute path. Returns null when the
// server cannot locate a matching directory, so the caller can fail visibly
// instead of committing a bare leaf name that would resolve against the
// server process working directory.
// Resolve a browser-picked folder name (which exposes only the leaf name)
// to a server-side absolute path; null when the server cannot locate it,
// so the caller fails visibly instead of committing a bare leaf name.
async function resolveNativeName(name: string): Promise<string | null> {
try {
const res = await ToolsService.executeToolRaw(BuiltInTool.FILE_GLOB_SEARCH, {
@@ -318,7 +231,7 @@
}
function handleSubmit() {
const value = inputValue.trim();
const value = query.trim();
if (!value) {
closePicker();
return;
@@ -330,47 +243,33 @@
function handleKeydown(event: KeyboardEvent) {
if (event.key === KeyboardKey.ENTER) {
event.preventDefault();
// Commit the highlighted result, falling back to the raw input
// only when the query returned no matches.
if (hoveredIndex >= 0 && queryResults[hoveredIndex]) {
commit(queryResults[hoveredIndex]);
if (nav.hoveredIndex >= 0 && queryResults[nav.hoveredIndex]) {
commit(queryResults[nav.hoveredIndex]);
} else if (queryResults.length === 0) {
handleSubmit();
}
} else if (event.key === KeyboardKey.ARROW_DOWN) {
if (queryResults.length > 0) {
event.preventDefault();
hoveredIndex = (hoveredIndex + 1) % queryResults.length;
scrollTrigger++;
nav.move(1);
}
} else if (event.key === KeyboardKey.ARROW_UP) {
if (queryResults.length > 0) {
event.preventDefault();
hoveredIndex = hoveredIndex <= 0 ? queryResults.length - 1 : hoveredIndex - 1;
scrollTrigger++;
nav.move(-1);
}
}
}
function handleInputInput(value: string) {
hoveredIndex = -1;
if (value.trim().length > 0) {
runSearch(value);
}
}
function clearDirectory(event?: MouseEvent) {
// Stop the click from bubbling into the popover trigger and re-opening
// Stop the click from bubbling into the chip button and re-opening
// the picker on top of the now-cleared state.
event?.stopPropagation();
event?.preventDefault();
directory = null;
onChange?.(null);
closePicker();
}
// The chip is always visible; the X clears the directory (no-op when
// already empty).
function handleDismiss(event?: MouseEvent) {
event?.stopPropagation();
event?.preventDefault();
@@ -380,105 +279,104 @@
}
function handleOpenChange(open: boolean) {
isOpen = open;
if (open) {
// Seed the search field with the current path so the user can refine it
// (or hit Enter to confirm / clear via the X icon).
inputValue = directory ?? '';
hoveredIndex = -1;
queryResults = [];
searchError = null;
void toolsStore.resolveServerHome();
searchScope = homeBase ?? HOME_TILDE;
if (inputValue.trim()) runSearch(inputValue);
} else {
cancelSearch();
// bits-ui-initiated close (Escape on the content, outside-click,
// trigger toggle) - the only path that bypasses closePicker().
search.cancel();
// bits-ui-initiated close (Escape on the content, outside-click) -
// the only path that bypasses closePicker().
onClose?.();
}
}
// Tooltips only on wider viewports - hover surfaces get in the way on
// touch / narrow layouts. Mirrors the gate used in ActionIcon.
let innerWidth = $state(0);
const showTooltip = $derived(innerWidth > DEFAULT_MOBILE_BREAKPOINT);
</script>
<div
<button
type="button"
class={[
'justify-self-start flex min-w-0 w-auto items-center gap-1 mt-1.5 py-1 px-2 backdrop-blur-2xl rounded-md',
className,
isOpen && 'w-full'
className
]}
onclick={onOpen}
{disabled}
>
<Popover.Root bind:open={isOpen} onOpenChange={handleOpenChange}>
<Popover.Trigger {disabled} class="flex justify-start">
<ChatFormWorkingDirectoryChip
{directory}
{homeBase}
{disabled}
{showTooltip}
onClear={handleDismiss}
<ChatFormWorkingDirectoryChip
{directory}
{homeBase}
{disabled}
{showTooltip}
onClear={handleDismiss}
/>
</button>
<Popover.Root open={isOpen} onOpenChange={handleOpenChange}>
<Popover.Trigger
class="pointer-events-none absolute inset-0 opacity-0"
tabindex={-1}
aria-hidden="true"
>
<span class="sr-only">Open working directory picker</span>
</Popover.Trigger>
<Popover.Content
side="top"
align="start"
sideOffset={12}
{customAnchor}
preventScroll={false}
onkeydown={handleKeydown}
onOpenAutoFocus={(event) => event.preventDefault()}
onCloseAutoFocus={(event) => event.preventDefault()}
class="w-[var(--bits-popover-anchor-width)] max-w-none rounded-xl border-border/50 p-0 shadow-xl"
>
<div class="p-2 min-h-22 flex flex-col justify-between">
<SearchInput
bind:ref={searchInputRef}
bind:value={query}
placeholder="Choose working directory"
onClose={closePicker}
class="w-full"
/>
</Popover.Trigger>
<Popover.Content
side="top"
align="start"
sideOffset={4}
class="md:max-w-3xl w-[calc(100vw-1rem)] rounded-xl border-border/50 p-0 shadow-xl md:-translate-2!"
onkeydown={handleKeydown}
onOpenAutoFocus={(event) => event.preventDefault()}
>
<div class="p-2 min-h-28 flex flex-col justify-between">
<SearchInput
bind:ref={searchInputRef}
bind:value={inputValue}
placeholder="Choose working directory"
onInput={handleInputInput}
onClose={closePicker}
class="w-full"
{#if query.trim() && (search.isSearching || queryResults.length > 0 || searchError)}
<ChatFormWorkingDirectoryResultsList
results={queryResults}
hoveredIndex={nav.hoveredIndex}
isSearching={search.isSearching}
error={searchError}
rawQuery={query}
bind:container={listContainer}
onCommit={commit}
onHover={(index) => nav.setHover(index)}
/>
{/if}
{#if inputValue.trim() && (isSearching || queryResults.length > 0 || searchError)}
<ChatFormWorkingDirectoryResultsList
results={queryResults}
{hoveredIndex}
{isSearching}
error={searchError}
rawQuery={inputValue}
bind:container={listContainer}
onCommit={commit}
onHover={(index) => (hoveredIndex = index)}
/>
{/if}
{#if pickerSupported}
<button
type="button"
class="-mt-1 flex cursor-pointer items-center gap-2 rounded-sm px-2 py-1.5 text-sm outline-hidden select-none hover:bg-accent hover:text-accent-foreground"
onclick={browseNative}
>
<FolderOpen class="size-4 shrink-0 text-muted-foreground" />
<span>Browse</span>
</button>
{/if}
{#if pickerSupported}
<button
type="button"
class="-mt-1 flex cursor-pointer items-center gap-2 rounded-sm px-2 py-1.5 text-sm outline-hidden select-none hover:bg-accent hover:text-accent-foreground"
onclick={browseNative}
{#if homeBase}
<div class="-mx-2 my-2 h-px bg-border/20" aria-hidden="true"></div>
<span class="px-2 py-1.5 font-mono text-[10px]">
Searching in:
<span class="truncate text-muted-foreground/70" title={searchScope}
>{abbreviateHome(searchScope, homeBase)}</span
>
<FolderOpen class="size-4 shrink-0 text-muted-foreground" />
<span>Browse</span>
</button>
{/if}
{#if homeBase}
<div class="-mx-2 my-1 h-px bg-border/20" aria-hidden="true"></div>
<span class="px-2 py-2 font-mono text-[10px]">
Searching in:
<span class="truncate text-muted-foreground/70" title={searchScope}
>{abbreviateHome(searchScope, homeBase)}</span
>
</span>
{/if}
</div>
</Popover.Content>
</Popover.Root>
</div>
</span>
{/if}
</div>
</Popover.Content>
</Popover.Root>
<svelte:window bind:innerWidth />
@@ -1,6 +1,7 @@
<script lang="ts">
import { Folder, X } from '@lucide/svelte';
import { abbreviateWorkingDir } from '$lib/utils';
import { SET_WORKING_DIRECTORY_LABEL } from '$lib/constants';
import * as Tooltip from '$lib/components/ui/tooltip';
import { ActionIcon } from '$lib/components/app/actions';
@@ -21,7 +22,7 @@
}: Props = $props();
const displayLabel = $derived(
directory ? abbreviateWorkingDir(directory, homeBase) : 'Select working directory'
directory ? abbreviateWorkingDir(directory, homeBase) : SET_WORKING_DIRECTORY_LABEL
);
// Full path surface: hover the abbreviated label to recall the exact directory.
const displayLabelTitle = $derived(directory ?? '');
@@ -183,8 +183,8 @@
<ChatMessageAssistantProcessingInfo {modelLoadingText} {processingState} position="bottom" />
{/if}
<div class="info my-6 grid gap-4 tabular-nums">
{#if displayedModel}
{#if displayedModel}
<div class="info my-6 grid gap-4 tabular-nums">
<div class="inline-flex flex-wrap items-start gap-2 text-xs text-muted-foreground">
<ChatMessageAssistantModel
{displayedModel}
@@ -200,8 +200,8 @@
showMessageStats={currentConfig.showMessageStats}
/>
</div>
{/if}
</div>
</div>
{/if}
{#if message.timestamp && !editCtx.isEditing}
<ChatMessageActionIcons
@@ -164,7 +164,7 @@
? `max-height: ${MAX_HEIGHT}px;`
: 'max-height: none;'}
>
{#if currentConfig.renderUserContentAsMarkdown}
{#if !currentConfig.renderContentAsRawText}
<div bind:this={messageElement} class={isExpanded ? 'cursor-text' : ''}>
<MarkdownContent class="markdown-system-content" content={message.content} />
</div>
@@ -98,9 +98,10 @@
showSpinner || (toolUi?.icon ?? null) || !mcpServerFavicon ? null : mcpServerFavicon
);
// No subtitle while the call is in flight - the spinner already
// signals activity; only terminal states get a pill.
function subtitleFor(errorMessage?: string): string | undefined {
if (extraLiveStreaming) return 'streaming...';
if (showSpinner) return 'executing...';
if (showSpinner) return undefined;
if (errorMessage) return 'failed';
if (isStreamingCall && !isStreaming) return 'incomplete';
return undefined;
@@ -63,7 +63,7 @@
data-multiline={isMultiline ? '' : undefined}
style="{maxHeightStyle} overflow-wrap: anywhere; word-break: break-word;"
>
{#if renderMarkdown && currentConfig.renderUserContentAsMarkdown}
{#if renderMarkdown && !currentConfig.renderContentAsRawText}
<div bind:this={messageElement}>
<MarkdownContent class="markdown-user-content" {content} />
</div>
@@ -41,7 +41,6 @@
let expandedStates: Record<number, boolean> = $state({});
const renderThinkingAsMarkdown = $derived(config().renderThinkingAsMarkdown as boolean);
const showThoughtInProgress = $derived(Boolean(config().showThoughtInProgress));
const alwaysShowToolCallContent = $derived(Boolean(config().alwaysShowToolCallContent));
const showMessageStats = $derived(Boolean(config().showMessageStats));
@@ -186,7 +185,6 @@
{section}
open={isExpanded(index, section)}
{isStreaming}
{renderThinkingAsMarkdown}
{hasReasoningError}
attachments={message?.extra}
onToggle={() => toggleExpanded(index, section)}
@@ -3,6 +3,7 @@
import { CollapsibleContentBlock, MarkdownContent } from '$lib/components/app';
import { AgenticSectionType } from '$lib/enums';
import { REASONING_SCROLL_AT_BOTTOM_THRESHOLD_PX } from '$lib/constants/auto-scroll';
import { config } from '$lib/stores/settings.svelte';
import type { DatabaseMessageExtra } from '$lib/types';
import type { AgenticSection } from '$lib/utils';
@@ -10,7 +11,6 @@
section: AgenticSection;
open: boolean;
isStreaming: boolean;
renderThinkingAsMarkdown: boolean;
hasReasoningError?: boolean;
attachments?: DatabaseMessageExtra[];
onToggle?: () => void;
@@ -20,12 +20,13 @@
section,
open,
isStreaming,
renderThinkingAsMarkdown,
hasReasoningError = false,
attachments,
onToggle
}: Props = $props();
const currentConfig = config();
const REASONING_HEADER = 'Reasoning';
const REASONING_HEADER_PENDING = 'Reasoning...';
const REASONING_SUBTITLE_ERROR = 'Error';
@@ -128,7 +129,7 @@
class:is-streaming={isPending}
onscroll={handleScrollEvent}
>
{#if renderThinkingAsMarkdown}
{#if !currentConfig.renderContentAsRawText}
<MarkdownContent content={section.content} class="text-muted-foreground" {attachments} />
{:else}
<div
+19 -26
View File
@@ -266,9 +266,9 @@ export { default as ChatFormFileInputInvisible } from './ChatForm/ChatFormFileIn
export { default as ChatFormMcpResourcesList } from './ChatForm/ChatFormMcpResourcesList.svelte';
/**
* Auto-resizing textarea with IME composition support. Automatically adjusts
* height based on content. Handles IME input correctly (waits for composition
* end before processing Enter key). Exposes focus() and resetHeight() methods.
* Auto-resizing textarea with IME composition support. Mention links stay
* plain markdown text in the input; the chip rendering happens in the
* message view via the rehype file-badge plugin.
*/
export { default as ChatFormTextarea } from './ChatForm/ChatFormTextarea.svelte';
@@ -351,14 +351,14 @@ export { default as ChatFormPickerPopover } from './ChatForm/ChatFormPickers/Cha
* Generic scrollable list for picker popovers. Provides search input,
* scroll-into-view for keyboard navigation, loading skeletons, empty state,
* and optional footer. Uses Svelte 5 snippets for item/skeleton/footer rendering.
* Shared by ChatFormPickerMcpPrompts and ChatFormPickerMcpResources.
* Shared by ChatFormPickerMcpPrompts and ChatFormMentionPicker.
*/
export { default as ChatFormPickerList } from './ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerList.svelte';
/**
* Generic button wrapper for picker list items. Provides consistent styling,
* hover/selected states, and data-picker-index attribute for scroll-into-view.
* Shared by ChatFormPickerMcpPrompts and ChatFormPickerMcpResources.
* Shared by ChatFormPickerMcpPrompts and ChatFormMentionPicker.
*/
export { default as ChatFormPickerListItem } from './ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerListItem.svelte';
@@ -376,30 +376,23 @@ export { default as ChatFormPickerItemHeader } from './ChatForm/ChatFormPickers/
export { default as ChatFormPickerListItemSkeleton } from './ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerListItemSkeleton.svelte';
/**
* **ChatFormPickerMcpResources** - MCP resource selection interface
*
* Floating picker for browsing and attaching MCP Server Resources.
* Triggered by typing `@` in the chat input.
* Loads resources from connected MCP servers and allows users to attach them to the chat context.
*
* **Features:**
* - Search/filter resources by name, title, description, or URI across all connected servers
* - Keyboard navigation (/ to navigate, Enter to select, Esc to close)
* - Shows attached state for already-attached resources
* - Loading states with skeleton placeholders
* - Server information header per resource for visual identification
*
* **Exported API:**
* - `handleKeydown(event): boolean` - Process keyboard events, returns true if handled
* `@`-triggered file/folder mention picker. Resolves `@<query>` in the chat
* input to a filesystem match via the server's `file_glob_search` built-in
* tool, scoped to the conversation cwd (or server home when unset).
* Selection splices a `[name](file:///<abs path>)` link into the input.
*/
export { default as ChatFormPickerMcpResources } from './ChatForm/ChatFormPickers/ChatFormPickerMcpResources.svelte';
export { default as ChatFormMentionPicker } from './ChatForm/ChatFormPickers/ChatFormMentionPicker.svelte';
/**
* **ChatFormPickers** - Chat input picker container
*
* Container component that hosts both MCP prompt and MCP resource pickers.
* Manages shared state, keyboard navigation, and coordination between the two
* picker interfaces. Used within ChatForm for `@`-triggered pickers.
* `/`-triggered slash-command picker. Lists the available slash commands
* (`/prompt`, `/cwd`, `/model`) filtered by the typed query; selection
* hands the command to the parent for dispatch.
*/
export { default as ChatFormCommandPicker } from './ChatForm/ChatFormPickers/ChatFormCommandPicker.svelte';
/**
* Hosts the chat-form pickers (slash-command, MCP prompt, file mention)
* and delegates keyboard events to the active one.
*/
export { default as ChatFormPickers } from './ChatForm/ChatFormPickers/ChatFormPickers.svelte';
@@ -15,6 +15,7 @@
import { SvelteMap } from 'svelte/reactivity';
import { rehypeRestoreTableHtml } from './plugins/rehype/table-html-restorer';
import { rehypeEnhanceLinks } from './plugins/rehype/enhance-links';
import { rehypeFileBadge } from './plugins/rehype/file-badge';
import { rehypeEnhanceCodeBlocks } from './plugins/rehype/enhance-code-blocks';
import { rehypeEnhanceMermaidBlocks } from './plugins/rehype/enhance-mermaid-blocks';
import { rehypeMermaidPre } from './plugins/rehype/mermaid-pre';
@@ -174,6 +175,7 @@
}) // Add syntax highlighting
.use(rehypeRestoreTableHtml) // Restore limited HTML (e.g., <br>, <ul>) inside Markdown tables
.use(rehypeEnhanceLinks) // Add target="_blank" to links
.use(rehypeFileBadge) // Render file:// anchors as inline badge chips
.use(rehypeMermaidPre) // Convert mermaid blocks to <pre class="mermaid">
.use(rehypeSvgPre) // Convert svg blocks to <pre class="svg-block">
.use(rehypeEnhanceCodeBlocks) // Wrap code blocks with header and actions
@@ -0,0 +1,100 @@
/**
* Rehype plugin that rewrites `file://` markdown anchors into the inline
* @-mention chip, reusing the visual contract from
* `$lib/constants/mention-badge`.
*
* The chip is presentational: `file://` navigation is blocked from
* http(s) pages, so the anchor becomes a plain `<span>` (no link role,
* no tab stop); the full path stays available on `title`.
*/
import { decodeFileLinkPath, getMentionBadgeIconPaths, getMentionBadgeLabel } from '$lib/utils';
import {
FILE_URI_PREFIX,
MENTION_BADGE_CLASSNAME,
MENTION_BADGE_ICON_CLASSNAME,
MENTION_BADGE_SVG_ATTRIBUTES,
PATH_SEPARATOR,
SETTINGS_KEYS
} from '$lib/constants';
import { settingsStore } from '$lib/stores/settings.svelte';
import { toolsStore } from '$lib/stores/tools.svelte';
import type { Plugin } from 'unified';
import type { Root, Element } from 'hast';
import { visit } from 'unist-util-visit';
// Trailing path separators mark a directory and are kept out of the label.
const TRAILING_SEPARATOR_REGEX = /\/+$/;
function decodeHrefPath(href: string): string {
const stripped = href.startsWith(FILE_URI_PREFIX) ? href.slice(FILE_URI_PREFIX.length) : href;
return decodeFileLinkPath(stripped);
}
function labelFromFileUrl(href: string): string {
const decoded = decodeHrefPath(href);
const trimmed = decoded.replace(TRAILING_SEPARATOR_REGEX, '');
const slash = trimmed.lastIndexOf(PATH_SEPARATOR);
return slash === -1 ? trimmed : trimmed.slice(slash + 1);
}
// A trailing `/` in the target marks a directory and selects the folder
// icon, matching the convention the mention picker inserts with.
function iconElement(href: string): Element {
return {
type: 'element',
tagName: 'svg',
properties: {
...MENTION_BADGE_SVG_ATTRIBUTES,
className: MENTION_BADGE_ICON_CLASSNAME.split(' ').filter(Boolean)
},
children: getMentionBadgeIconPaths(href).map((d) => ({
type: 'element',
tagName: 'path',
properties: { d },
children: []
}))
};
}
export const rehypeFileBadge: Plugin<[], Root> = () => {
return (tree: Root) => {
visit(tree, 'element', (node: Element) => {
if (node.tagName !== 'a') return;
const props = node.properties ?? {};
const href = typeof props.href === 'string' ? props.href : null;
if (!href || !href.startsWith(FILE_URI_PREFIX)) return;
const label = labelFromFileUrl(href);
const titleAttr = typeof props.title === 'string' ? props.title : href;
const decodedPath = decodeHrefPath(href);
node.tagName = 'span';
node.properties = {
className: MENTION_BADGE_CLASSNAME.split(' ').filter(Boolean),
title: titleAttr.startsWith(FILE_URI_PREFIX) ? decodedPath : titleAttr
};
node.children = [
iconElement(href),
{
type: 'element',
tagName: 'span',
properties: { className: ['shrink-0', 'truncate'] },
children: [
{
type: 'text',
value: getMentionBadgeLabel(
label,
decodedPath,
settingsStore.getConfig(SETTINGS_KEYS.SHOW_FULL_PATH_IN_MENTIONS),
toolsStore.serverHome
)
}
]
}
];
});
};
};
@@ -1,5 +1,6 @@
<script lang="ts">
import { ICON_CLASS_DEFAULT } from '$lib/constants/css-classes';
import { URL_PARAMS } from '$lib/constants';
import * as AlertDialog from '$lib/components/ui/alert-dialog';
import { AlertTriangle, ArrowRight } from '@lucide/svelte';
import { goto } from '$app/navigation';
@@ -22,7 +23,7 @@
function handleSelectModel(model: string) {
// Build URL with selected model, preserving other params
const url = new URL(page.url);
url.searchParams.set('model', model);
url.searchParams.set(URL_PARAMS.MODEL, model);
handleOpenChange(false);
goto(url.toString());
@@ -0,0 +1,25 @@
<script lang="ts">
import { highlightMatch } from '$lib/utils';
interface Props {
text: string;
query: string;
matchClass?: string;
}
let {
text,
query,
matchClass = 'rounded bg-yellow-200/60 px-0.5 text-foreground dark:bg-yellow-500/30'
}: Props = $props();
let segments = $derived(highlightMatch(text, query));
</script>
{#each segments as seg, i (i)}
{#if seg.match}
<mark class={matchClass}>{seg.text}</mark>
{:else}
{seg.text}
{/if}
{/each}
@@ -42,3 +42,11 @@ export { default as KeyValuePairs } from './KeyValuePairs.svelte';
* Supports placeholder, autofocus, and change callbacks.
*/
export { default as SearchInput } from './SearchInput.svelte';
/**
* **HighlightedMatch** - Substring-match text highlight
*
* Renders `text` with each case-insensitive occurrence of `query` wrapped
* in `<mark>`.
*/
export { default as HighlightedMatch } from './HighlightedMatch.svelte';
@@ -1,8 +1,9 @@
<script lang="ts">
import { ChevronDown, Loader2, Package } from '@lucide/svelte';
import { ChevronDown, Loader2 } from '@lucide/svelte';
import * as DropdownMenu from '$lib/components/ui/dropdown-menu';
import * as Tooltip from '$lib/components/ui/tooltip';
import { KeyboardKey, ServerModelStatus } from '$lib/enums';
import { MODEL_SELECTOR_ICON } from '$lib/constants';
import { useModelsSelector } from '$lib/hooks/use-models-selector.svelte';
import { modelsStore, routerModels } from '$lib/stores/models.svelte';
import { modelLoadFraction } from '$lib/utils';
@@ -35,7 +36,7 @@
}: Props = $props();
let isOpen = $state(false);
let highlightedIndex = $state<number>(-1);
let highlightedId = $state<string | null>(null);
const ms = useModelsSelector({
currentModel: () => currentModel,
@@ -43,15 +44,77 @@
onModelChange: () => onModelChange,
onOpenChange: (open) => {
isOpen = open;
highlightedIndex = -1;
highlightedId = null;
}
});
$effect(() => {
void ms.searchTerm;
highlightedIndex = -1;
highlightedId = null;
});
// Focus the dropdown's search box without scrolling the page. bits-ui
// auto-focuses the opened content by default, which can yank the page
// scroll; we prevent that on the Content and refocus the search here.
$effect(() => {
if (!isOpen) return;
requestAnimationFrame(() => {
const search = document.querySelector<HTMLElement>(
'[data-slot="dropdown-menu-content"] input'
);
search?.focus({ preventScroll: true });
});
});
// Keyboard navigation follows the on-screen row order, not the flat option list order.
let visualOrder = $derived.by(() => {
const order: string[] = [];
for (const item of ms.groupedFilteredOptions.loaded) order.push(item.option.id);
for (const item of ms.groupedFilteredOptions.favorites) order.push(item.option.id);
for (const group of ms.groupedFilteredOptions.available) {
for (const item of group.items) order.push(item.option.id);
}
return order;
});
let highlightedIndex = $derived(highlightedId ? visualOrder.indexOf(highlightedId) : -1);
function moveHighlight(direction: 1 | -1) {
const len = visualOrder.length;
if (len === 0) {
highlightedId = null;
return;
}
let index = highlightedIndex;
if (index === -1) {
index = direction === 1 ? 0 : len - 1;
} else {
index = (index + direction + len) % len;
}
highlightedId = visualOrder[index];
}
// Alt+Enter only unloads and keeps the dropdown open.
async function handleModelKeyAction(modelId: string, unload: boolean) {
if (!unload) {
void ms.handleSelect(modelId);
return;
}
const model = routerModels().find((m) => m.id === modelId);
const status = model?.status?.value as ServerModelStatus | undefined;
if (status === ServerModelStatus.LOADING) return;
await modelsStore.unloadModel(modelId);
}
export function open() {
ms.handleOpenChange(true);
}
@@ -61,33 +124,17 @@
if (event.key === KeyboardKey.ARROW_DOWN) {
event.preventDefault();
if (ms.filteredOptions.length === 0) return;
if (highlightedIndex === -1 || highlightedIndex === ms.filteredOptions.length - 1) {
highlightedIndex = 0;
} else {
highlightedIndex += 1;
}
moveHighlight(1);
} else if (event.key === KeyboardKey.ARROW_UP) {
event.preventDefault();
if (ms.filteredOptions.length === 0) return;
if (highlightedIndex === -1 || highlightedIndex === 0) {
highlightedIndex = ms.filteredOptions.length - 1;
} else {
highlightedIndex -= 1;
}
moveHighlight(-1);
} else if (event.key === KeyboardKey.ENTER) {
event.preventDefault();
if (highlightedIndex >= 0 && highlightedIndex < ms.filteredOptions.length) {
const option = ms.filteredOptions[highlightedIndex];
ms.handleSelect(option.id);
} else if (ms.filteredOptions.length > 0) {
highlightedIndex = 0;
if (highlightedId) {
void handleModelKeyAction(highlightedId, event.altKey);
} else if (visualOrder.length > 0) {
highlightedId = visualOrder[0];
}
}
}
@@ -109,7 +156,7 @@
]}
style="max-width: min(calc(100cqw - 10rem), 20rem)"
>
<Package class="h-3.5 w-3.5 shrink-0" />
<MODEL_SELECTOR_ICON class="h-3.5 w-3.5 shrink-0" />
</span>
{:else}
<p class="text-xs text-muted-foreground">No models available.</p>
@@ -150,7 +197,7 @@
]}
disabled={disabled || ms.updating}
>
<Package class="h-3.5 w-3.5 shrink-0" />
<MODEL_SELECTOR_ICON class="h-3.5 w-3.5 shrink-0" />
{#if selectedOption}
<ModelId
@@ -186,6 +233,7 @@
<DropdownMenu.Content
align="end"
class="w-full max-w-[100vw] pt-0 sm:w-max sm:max-w-[calc(100vw-2rem)]"
onOpenAutoFocus={(event) => event.preventDefault()}
>
<DropdownMenuSearchable
searchValue={ms.searchTerm}
@@ -217,9 +265,9 @@
{/if}
{#snippet modelOption(item: ModelItem, hideOrgName: boolean)}
{@const { option, flatIndex } = item}
{@const { option } = item}
{@const isSelected = currentModel === option.model || ms.activeId === option.id}
{@const isHighlighted = flatIndex === highlightedIndex}
{@const isHighlighted = option.id === highlightedId}
{@const isFav = ms.isFavorite(option.model)}
<ModelsSelectorOption
@@ -230,11 +278,11 @@
{hideOrgName}
onSelect={ms.handleSelect}
onInfoClick={ms.handleInfoClick}
onMouseEnter={() => (highlightedIndex = flatIndex)}
onMouseEnter={() => (highlightedId = option.id)}
onKeyDown={(event) => {
if (event.key === KeyboardKey.ENTER || event.key === KeyboardKey.SPACE) {
event.preventDefault();
ms.handleSelect(option.id);
void handleModelKeyAction(option.id, event.altKey);
}
}}
/>
@@ -275,7 +323,7 @@
onclick={() => ms.handleOpenChange(true)}
disabled={disabled || ms.updating}
>
<Package class="h-3.5 w-3.5 shrink-0" />
<MODEL_SELECTOR_ICON class="h-3.5 w-3.5 shrink-0" />
{#if selectedOption}
<ModelId
@@ -62,9 +62,10 @@
<div
class={[
'group relative flex w-full items-center gap-2 rounded-sm p-2 text-left text-sm transition focus:outline-none',
'cursor-pointer hover:bg-muted focus:bg-muted',
(isSelected || isHighlighted) && 'bg-accent text-accent-foreground',
!(isSelected || isHighlighted) && 'hover:bg-accent hover:text-accent-foreground',
'cursor-pointer',
isSelected && 'bg-accent/50 text-accent-foreground',
isHighlighted && 'bg-accent',
!isSelected && !isHighlighted && 'hover:bg-muted',
isLoaded ? 'text-popover-foreground' : 'text-muted-foreground'
]}
role="option"
@@ -98,7 +98,12 @@
const numValue = Number(processedConfig[field]);
if (!isNaN(numValue)) {
if ((POSITIVE_INTEGER_FIELDS as readonly string[]).includes(field)) {
processedConfig[field] = Math.max(1, Math.round(numValue));
const entryByMinMax = SETTINGS_CHAT_SECTIONS.flatMap(
(section) => section.fields ?? []
).find((entry) => entry.key === field);
const lo = entryByMinMax?.min ?? 1;
const hi = entryByMinMax?.max ?? Number.POSITIVE_INFINITY;
processedConfig[field] = Math.max(lo, Math.min(hi, Math.round(numValue)));
} else {
processedConfig[field] = numValue;
}
@@ -83,12 +83,18 @@
<Input
id={field.key}
type={field.isPositiveInteger ? 'number' : 'text'}
{...field.isPositiveInteger ? { min: '1', step: '1' } : {}}
{...field.isPositiveInteger
? {
min: String(field.min ?? 1),
step: '1',
...(field.max != null ? { max: String(field.max) } : {})
}
: {}}
value={currentValue}
oninput={(e) => onConfigChange(field.key, e.currentTarget.value)}
placeholder={currentModelParams[field.key] != null
? `Default: ${normalizeFloatingPoint(currentModelParams[field.key])}`
: ''}
: (field.placeholder ?? '')}
class="w-full {isCustomRealTime ? 'pr-8' : ''}"
/>
{#if isCustomRealTime}
@@ -0,0 +1,44 @@
import { SET_WORKING_DIRECTORY_LABEL } from '$lib/constants/working-directory';
import { ChatFormCommandAction } from '$lib/enums';
import type { ChatFormCommand } from '$lib/types';
interface ChatCommandsOptions {
/** Gates `/model`. */
showModelSelector: boolean;
/** Gates `/prompt`. */
hasPrompts: () => boolean;
/** Gates `/cwd`. */
hasBuiltinTools: () => boolean;
}
/**
* The slash commands surfaced by the `/` command picker, in display order.
*
* Availability is supplied as predicates rather than store imports: this
* module is re-exported through the `$lib/constants` barrel, and importing
* stores at module load would create a circular dependency (the stores
* themselves import from `$lib/constants`).
*/
export function getChatCommands(options: ChatCommandsOptions): ChatFormCommand[] {
return [
{
name: 'prompt',
description: 'Insert an MCP prompt',
action: ChatFormCommandAction.PROMPT,
disabled: !options.hasPrompts()
},
{
name: 'cwd',
description: SET_WORKING_DIRECTORY_LABEL,
keywords: ['current working directory'],
action: ChatFormCommandAction.CWD,
disabled: !options.hasBuiltinTools()
},
{
name: 'model',
description: 'Select model',
action: ChatFormCommandAction.MODEL,
disabled: !options.showModelSelector
}
];
}
-1
View File
@@ -2,5 +2,4 @@ export const INITIAL_FILE_SIZE = 0;
export const PROMPT_CONTENT_SEPARATOR = '\n\n';
export const CLIPBOARD_CONTENT_QUOTE_PREFIX = '"';
export const PROMPT_TRIGGER_PREFIX = '/';
export const RESOURCE_TRIGGER_PREFIX = '@';
export const NEW_CHAT_DRAFT_KEY = '__new_chat__';
@@ -19,6 +19,9 @@ export const PANEL_CLASSES = `
export const CHAT_FORM_POPOVER_MAX_HEIGHT = 'max-h-80';
export const DIALOG_SUBMENU_CONTENT = 'w-60';
/** Selects the chat-form input to restore focus after model actions. */
export const CHAT_INPUT_FOCUS_SELECTOR = '[data-slot="input-area"] textarea';
/** Default Tailwind size class for inline icon components (lucide, etc.). */
export const ICON_CLASS_DEFAULT = 'h-4 w-4';
+2
View File
@@ -18,6 +18,7 @@ export * from './binary-detection';
export * from './built-in-tools';
export * from './cache';
export * from './chat-form';
export * from './chat-commands';
export * from './cli-flags';
export * from './code-blocks';
export * from './icons';
@@ -39,6 +40,7 @@ export * from './max-bundle-size';
export * from './mcp';
export * from './mcp-form';
export * from './mcp-resource';
export * from './mention-badge';
export * from './message-export';
export * from './path-display';
export * from './model-id';
@@ -0,0 +1,39 @@
/**
* Visual contract for message @-mention badges. Svelte cannot be mounted
* from a hast tree, so the rehype file-badge plugin emits the shared class
* string below; keeping it here as a literal lets Tailwind's source
* scanner generate the utility classes.
*/
export const MENTION_BADGE_CLASSNAME =
'inline-flex w-fit shrink-0 items-center gap-1 whitespace-nowrap rounded-md border border-border/50 bg-foreground/5 px-1.5 py-0.5 text-xs font-mono text-foreground hover:bg-foreground/10 dark:bg-foreground/10 dark:text-secondary-foreground';
export const MENTION_BADGE_ICON_CLASSNAME = 'h-3 w-3 shrink-0';
/**
* SVG attributes shared by the hast-built badge icons; the rehype plugin
* spreads them onto the `<svg>` `properties`.
*/
export const MENTION_BADGE_SVG_ATTRIBUTES: Readonly<Record<string, string>> = {
xmlns: 'http://www.w3.org/2000/svg',
viewBox: '0 0 24 24',
fill: 'none',
stroke: 'currentColor',
'stroke-width': '2',
'stroke-linecap': 'round',
'stroke-linejoin': 'round',
'aria-hidden': 'true'
};
/**
* SVG path strings for the badge's inline icon; each entry becomes one
* `<path>` child of the wrapper `<svg>`. Paths match `lucide-svelte`'s
* current `File` and `Folder` glyphs.
*/
export const MENTION_BADGE_FILE_ICON_PATHS: readonly string[] = [
'M6 22a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h8a2.4 2.4 0 0 1 1.704.706l3.588 3.588A2.4 2.4 0 0 1 20 8v12a2 2 0 0 1-2 2z',
'M14 2v5a1 1 0 0 0 1 1h5'
];
export const MENTION_BADGE_FOLDER_ICON_PATHS: readonly string[] = [
'M20 20a2 2 0 0 0 2-2V8a2 2 0 0 0-2-2h-7.9a2 2 0 0 1-1.69-.9L9.6 3.9A2 2 0 0 0 7.93 3H4a2 2 0 0 0-2 2v13a2 2 0 0 0 2 2Z'
];
+12 -2
View File
@@ -1,4 +1,14 @@
export const NEW_CHAT_PARAM = 'new_chat';
/** Query params the chat routes read from the URL. */
export const URL_PARAMS = {
/** Prompt to send on arrival. */
QUERY: 'q',
/** Model to select. */
MODEL: 'model',
/** Load the selected model instead of waiting for the first message. */
LOAD: 'load',
/** Start a new chat. */
NEW_CHAT: 'new_chat'
} as const;
/** Settings section slugs — used for routes and navigation. */
export const SETTINGS_SECTION_SLUGS = {
@@ -16,7 +26,7 @@ export const ROUTES = {
/** Root — start of the app. */
START: '#/',
/** New chat — root with new chat query param. */
NEW_CHAT: `?${NEW_CHAT_PARAM}=true#/`,
NEW_CHAT: `?${URL_PARAMS.NEW_CHAT}=true#/`,
/** Chat base — for dynamic chat URLs use RouterService. */
CHAT: '#/chat',
/** MCP servers. */
+3 -2
View File
@@ -23,7 +23,7 @@ export const SETTINGS_KEYS = {
SHOW_AGENTIC_TURN_STATS: 'showAgenticTurnStats',
SHOW_THOUGHT_IN_PROGRESS: 'showThoughtInProgress',
AUTO_MIC_ON_EMPTY: 'autoMicOnEmpty',
RENDER_USER_CONTENT_AS_MARKDOWN: 'renderUserContentAsMarkdown',
RENDER_CONTENT_AS_RAW_TEXT: 'renderContentAsRawText',
DISABLE_AUTO_SCROLL: 'disableAutoScroll',
ALWAYS_SHOW_SIDEBAR_ON_DESKTOP: 'alwaysShowSidebarOnDesktop',
FULL_HEIGHT_CODE_BLOCKS: 'fullHeightCodeBlocks',
@@ -31,8 +31,9 @@ export const SETTINGS_KEYS = {
SHOW_MODEL_QUANTIZATION: 'showModelQuantization',
SHOW_MODEL_TAGS: 'showModelTags',
SHOW_BUILD_VERSION: 'showBuildVersion',
SHOW_FULL_PATH_IN_MENTIONS: 'showFullPathInMentions',
SHOW_SYSTEM_MESSAGE: 'showSystemMessage',
RENDER_THINKING_AS_MARKDOWN: 'renderThinkingAsMarkdown',
MENTION_SEARCH_MAX_DEPTH: 'mentionSearchMaxDepth',
// Sampling
TEMPERATURE: 'temperature',
DYNATEMP_RANGE: 'dynatemp_range',
+32 -12
View File
@@ -23,7 +23,12 @@ import type {
SettingsSectionEntry,
SettingsSection
} from '$lib/types';
import { CLI_FLAGS, DEFAULT_MCP_CONFIG } from '$lib/constants';
import { CLI_FLAGS } from './cli-flags';
import { DEFAULT_MCP_CONFIG } from './mcp';
import {
FILE_GLOB_SEARCH_PICKERS_MAX_SEARCH_DEPTH,
FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH
} from './working-directory';
import { SETTINGS_KEYS } from './settings-keys';
import { ROUTES, SETTINGS_SECTION_SLUGS } from './routes';
import { TITLE_GENERATION } from './title-generation';
@@ -228,21 +233,13 @@ const SETTINGS_REGISTRY: Record<string, SettingsSectionEntry> = {
section: SETTINGS_SECTION_SLUGS.DISPLAY
},
{
key: SETTINGS_KEYS.RENDER_USER_CONTENT_AS_MARKDOWN,
label: 'Render user content as Markdown',
help: 'Render user messages using markdown formatting in the chat.',
key: SETTINGS_KEYS.RENDER_CONTENT_AS_RAW_TEXT,
label: 'Render content as raw text',
help: 'Display user, system and thinking content as plain text instead of formatted Markdown. Markdown is the default so that @-mention badges render in sent messages.',
defaultValue: false,
type: SettingsFieldType.CHECKBOX,
section: SETTINGS_SECTION_SLUGS.DISPLAY
},
{
key: SETTINGS_KEYS.RENDER_THINKING_AS_MARKDOWN,
label: 'Render thinking as Markdown',
help: 'Render the reasoning/thinking block content as formatted Markdown instead of plain text.',
defaultValue: true,
type: SettingsFieldType.CHECKBOX,
section: SETTINGS_SECTION_SLUGS.DISPLAY
},
{
key: SETTINGS_KEYS.FULL_HEIGHT_CODE_BLOCKS,
label: 'Use full height code blocks',
@@ -298,6 +295,14 @@ const SETTINGS_REGISTRY: Record<string, SettingsSectionEntry> = {
defaultValue: false,
type: SettingsFieldType.CHECKBOX,
section: SETTINGS_SECTION_SLUGS.DISPLAY
},
{
key: SETTINGS_KEYS.SHOW_FULL_PATH_IN_MENTIONS,
label: 'Show full path in mentions',
help: 'Display the full file system path inside file and folder @-mention badges instead of just the file or folder name.',
defaultValue: false,
type: SettingsFieldType.CHECKBOX,
section: SETTINGS_SECTION_SLUGS.DISPLAY
}
]
},
@@ -555,6 +560,18 @@ const SETTINGS_REGISTRY: Record<string, SettingsSectionEntry> = {
type: SettingsFieldType.INPUT,
section: SETTINGS_SECTION_SLUGS.AGENTIC,
isPositiveInteger: true
},
{
key: SETTINGS_KEYS.MENTION_SEARCH_MAX_DEPTH,
label: 'Mention search depth',
help: 'How many directory levels below the working directory the @-mention file search descends. Larger values surface deeply nested files but take longer on large trees.',
defaultValue: FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH,
placeholder: `${FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH}`,
min: 1,
max: FILE_GLOB_SEARCH_PICKERS_MAX_SEARCH_DEPTH,
type: SettingsFieldType.INPUT,
section: SETTINGS_SECTION_SLUGS.AGENTIC,
isPositiveInteger: true
}
]
},
@@ -699,6 +716,9 @@ export const SETTINGS_CHAT_SECTIONS: SettingsSection[] = [
type: s.type,
isExperimental: s.isExperimental,
isPositiveInteger: s.isPositiveInteger,
placeholder: s.placeholder,
min: s.min,
max: s.max,
dependsOn: s.dependsOn,
help: s.help,
options: s.options,
+4 -1
View File
@@ -1,4 +1,4 @@
import { Search, Settings, SquarePen } from '@lucide/svelte';
import { Package, Search, Settings, SquarePen } from '@lucide/svelte';
import McpLogo from '$lib/components/app/mcp/McpLogo.svelte';
import type { Component } from 'svelte';
import { ROUTES } from './routes';
@@ -6,6 +6,9 @@ import { ROUTES } from './routes';
export const FORK_TREE_DEPTH_PADDING = 8;
export const SYSTEM_MESSAGE_PLACEHOLDER = 'System message';
/** Icon used for the model selector and the `/model` slash command. */
export const MODEL_SELECTOR_ICON = Package;
export const ICON_STRIP_TRANSITION_DURATION = 150;
export const ICON_STRIP_TRANSITION_DELAY_MULTIPLIER = 50;
@@ -8,6 +8,9 @@
export const GLOB_WILDCARD = '*';
/** Label shown for the working-directory picker / `/cwd` slash command. */
export const SET_WORKING_DIRECTORY_LABEL = 'Set working directory';
/** Character that starts and ends a glob character-class fragment. */
export const GLOB_RANGE_OPEN = '[';
export const GLOB_RANGE_CLOSE = ']';
@@ -38,3 +41,9 @@ export const PATH_NAV_MAX_DEPTH = 1;
// Native folder-picker resolution searches a shallow, bounded window.
export const NATIVE_MAX_DEPTH = 4;
export const NATIVE_LIMIT = 20;
/** Upper bound the mention search depth setting accepts. The server itself imposes no depth cap (0 = unlimited); this is a UI sanity bound. */
export const FILE_GLOB_SEARCH_PICKERS_MAX_SEARCH_DEPTH = 32;
/** Depth the pickers fall back to when the user setting is invalid. */
export const FILE_GLOB_SEARCH_PICKERS_DEFAULT_SEARCH_DEPTH = 10;
+11
View File
@@ -78,3 +78,14 @@ export enum PdfViewMode {
TEXT = 'text',
PAGES = 'pages'
}
export enum ChatFormCommandAction {
PROMPT = 'prompt',
CWD = 'cwd',
MODEL = 'model'
}
export enum FileMentionEntryType {
FILE = 'file',
DIRECTORY = 'directory'
}
+3 -1
View File
@@ -24,7 +24,9 @@ export {
MessageRole,
MessageType,
PdfViewMode,
ReasoningFormat
ReasoningFormat,
ChatFormCommandAction,
FileMentionEntryType
} from './chat.enums';
export { SessionRecordType } from './conversation-import.enums';
+1
View File
@@ -8,6 +8,7 @@ export enum KeyboardKey {
ARROW_DOWN = 'ArrowDown',
ARROW_LEFT = 'ArrowLeft',
ARROW_RIGHT = 'ArrowRight',
BACKSPACE = 'Backspace',
TAB = 'Tab',
B_LOWER = 'b',
D_LOWER = 'd',
@@ -0,0 +1,355 @@
import { getChatCommands, PROMPT_TRIGGER_PREFIX } from '$lib/constants';
import { ChatFormCommandAction, KeyboardKey } from '$lib/enums';
import type { ChatFormCommand } from '$lib/types';
import {
findCommandToken,
findMentionToken,
takeCommandDismissSnapshot,
takeMentionDismissSnapshot,
type CommandDismissSnapshot,
type MentionDismissSnapshot
} from '$lib/utils';
/** Dependencies injected as getters so the hook stays free of store circular imports. */
export interface UseChatFormPickersOptions {
getValue: () => string;
/** Also fires the form's onChange. */
setValue: (value: string) => void;
/** Undefined when unmounted. */
getCaretOffset: () => number | undefined;
setCaretOffset: (offset: number) => void;
focusInput: () => void;
/** Gates `/model`. */
getShowModelSelector: () => boolean;
/** Gates `/prompt`. */
hasPrompts: () => boolean;
/** Gates `/cwd`. */
hasBuiltinTools: () => boolean;
getCwd: () => string | null;
/** Mention search fallback scope. */
getServerHome: () => string | null;
openModelSelector: () => void;
/** Delegate a keydown to the mounted pickers component, if any. */
getPickersRef: () => { handleKeydown(event: KeyboardEvent): boolean } | undefined;
}
/**
* Chat-form picker state and the `/`+`@` routing that drives them.
* Owns open/query state, dismiss snapshots and slash-command dispatch;
* textarea/caret/attachment handling stays in the chat form.
*/
export function useChatFormPickers(opts: UseChatFormPickersOptions) {
let isCommandPickerOpen = $state(false);
let commandQuery = $state('');
let isPromptPickerOpen = $state(false);
let promptSearchQuery = $state('');
let isMentionPickerOpen = $state(false);
let mentionQuery = $state('');
let isWorkingDirectoryPickerOpen = $state(false);
let workingDirectoryQuery = $state('');
// Last dismissed `@`-mention token; while intact, the picker does not
// reopen, so an escaped `@<query>` stays literal until edited.
let mentionDismissedSnapshot: MentionDismissSnapshot | null = null;
// Same dismissal contract for the `/`-command token.
let commandDismissedSnapshot: CommandDismissSnapshot | null = null;
// Fall back to the server home so the picker still finds matches
// before a cwd is set.
const mentionScopePath = $derived(opts.getCwd() ?? opts.getServerHome() ?? null);
const availableCommands = $derived(
getChatCommands({
showModelSelector: opts.getShowModelSelector(),
hasPrompts: opts.hasPrompts,
hasBuiltinTools: opts.hasBuiltinTools
})
);
// Dispatch a slash command picked from the list: consume the token and
// open the target picker, seeding its search with `args`. Runs only on
// explicit selection (Enter/click), so the buffer is never cleared
// mid-typing.
function dispatchCommand(command: ChatFormCommand, args: string) {
isCommandPickerOpen = false;
commandQuery = '';
switch (command.action) {
case ChatFormCommandAction.PROMPT:
isWorkingDirectoryPickerOpen = false;
opts.setValue('');
isPromptPickerOpen = true;
promptSearchQuery = args.trim();
break;
case ChatFormCommandAction.CWD: {
// Keep `/cwd <args>` in the input so the search field and the
// token stay two-way bound; normalize partial tokens (`/cw foo`).
const trimmed = args.trim();
const newValue = `/cwd ${trimmed}`;
if (opts.getValue() !== newValue) {
opts.setValue(newValue);
queueMicrotask(() => opts.setCaretOffset(newValue.length));
}
workingDirectoryQuery = trimmed;
isWorkingDirectoryPickerOpen = true;
break;
}
case ChatFormCommandAction.MODEL:
isWorkingDirectoryPickerOpen = false;
opts.setValue('');
opts.openModelSelector();
break;
}
}
function handleInput() {
const value = opts.getValue();
const cursor = opts.getCaretOffset() ?? value.length;
if (value.startsWith(PROMPT_TRIGGER_PREFIX)) {
isMentionPickerOpen = false;
mentionQuery = '';
isPromptPickerOpen = false;
promptSearchQuery = '';
const token = findCommandToken(value);
if (!token) {
isCommandPickerOpen = false;
commandQuery = '';
return;
}
// While the `/cwd` picker is open the token doubles as its search
// field: keep the two in sync instead of re-dispatching.
if (isWorkingDirectoryPickerOpen) {
isCommandPickerOpen = false;
commandQuery = '';
if (token.name === 'cwd') {
workingDirectoryQuery = token.args.trim();
} else {
isWorkingDirectoryPickerOpen = false;
workingDirectoryQuery = '';
}
return;
}
// Dismissed token stays literal until it changes.
const isDismissedSticky =
commandDismissedSnapshot !== null &&
commandDismissedSnapshot.name === token.name &&
commandDismissedSnapshot.args === token.args;
if (isDismissedSticky) {
isCommandPickerOpen = false;
commandQuery = '';
return;
}
// Commands dispatch only on explicit selection (Enter/click),
// never mid-typing: `/model is broken` is prose until the user
// picks the command from the list.
if (availableCommands.length > 0) {
isCommandPickerOpen = true;
commandQuery = token.name;
} else {
isCommandPickerOpen = false;
commandQuery = '';
}
return;
}
isCommandPickerOpen = false;
commandQuery = '';
if (commandDismissedSnapshot !== null) {
commandDismissedSnapshot = null;
}
if (isWorkingDirectoryPickerOpen) {
isWorkingDirectoryPickerOpen = false;
}
const token = findMentionToken(value, cursor);
if (token) {
// Dismissed token stays literal: don't reopen until it changes.
const isDismissedSticky =
mentionDismissedSnapshot !== null &&
mentionDismissedSnapshot.start === token.start &&
mentionDismissedSnapshot.query === token.query;
if (!isDismissedSticky) {
// Only search once a char follows `@`; a bare `@` is a no-op
// (otherwise the picker flashes an empty hint on re-type).
if (token.query.length > 0) {
mentionDismissedSnapshot = null;
isMentionPickerOpen = true;
mentionQuery = token.query;
isPromptPickerOpen = false;
promptSearchQuery = '';
return;
}
}
}
isPromptPickerOpen = false;
promptSearchQuery = '';
isMentionPickerOpen = false;
mentionQuery = '';
// Token gone or changed: reset the snapshot so a fresh `@` reopens.
if (mentionDismissedSnapshot !== null && !token) {
mentionDismissedSnapshot = null;
}
}
function handleKeydown(event: KeyboardEvent): boolean {
if (opts.getPickersRef()?.handleKeydown(event)) {
return true;
}
if (event.key === KeyboardKey.ESCAPE && isPromptPickerOpen) {
isPromptPickerOpen = false;
promptSearchQuery = '';
return true;
}
return false;
}
function handleCommandSelect(command: ChatFormCommand) {
// Dispatch on the live token so typed args seed the target picker.
const token = findCommandToken(opts.getValue());
dispatchCommand(command, token?.args ?? '');
}
// Picker dismissed: snapshot the live token so it stays literal until
// deleted or retyped.
function handleCommandPickerClose() {
if (isCommandPickerOpen) {
commandDismissedSnapshot = takeCommandDismissSnapshot(opts.getValue());
}
isCommandPickerOpen = false;
commandQuery = '';
// Target picker manages its own focus: don't yank it back to the input.
if (!isPromptPickerOpen && !isMentionPickerOpen && !isWorkingDirectoryPickerOpen) {
opts.focusInput();
}
}
// Same dismissal snapshot for the mention token.
function handleMentionPickerClose() {
if (isMentionPickerOpen) {
const cursor = opts.getCaretOffset() ?? opts.getValue().length;
mentionDismissedSnapshot = takeMentionDismissSnapshot(opts.getValue(), cursor);
}
isMentionPickerOpen = false;
mentionQuery = '';
opts.focusInput();
}
function handlePromptPickerClose() {
isPromptPickerOpen = false;
promptSearchQuery = '';
opts.focusInput();
}
function handleWorkingDirectoryOpen() {
workingDirectoryQuery = opts.getCwd() ?? '';
isWorkingDirectoryPickerOpen = true;
}
function handleWorkingDirectoryClose() {
isWorkingDirectoryPickerOpen = false;
workingDirectoryQuery = '';
opts.focusInput();
}
// Two-way bind the text after `/cwd ` and the picker search input; the
// reverse direction is handled by handleInput.
$effect(() => {
if (!isWorkingDirectoryPickerOpen) return;
const value = opts.getValue();
const token = findCommandToken(value);
if (!token || token.name !== 'cwd') return;
const newValue = `/cwd ${workingDirectoryQuery}`;
if (newValue === value) return;
opts.setValue(newValue);
queueMicrotask(() => opts.setCaretOffset(newValue.length));
});
return {
get isCommandPickerOpen() {
return isCommandPickerOpen;
},
set isCommandPickerOpen(v: boolean) {
isCommandPickerOpen = v;
},
get commandQuery() {
return commandQuery;
},
set commandQuery(v: string) {
commandQuery = v;
},
get isPromptPickerOpen() {
return isPromptPickerOpen;
},
set isPromptPickerOpen(v: boolean) {
isPromptPickerOpen = v;
},
get promptSearchQuery() {
return promptSearchQuery;
},
set promptSearchQuery(v: string) {
promptSearchQuery = v;
},
get isMentionPickerOpen() {
return isMentionPickerOpen;
},
set isMentionPickerOpen(v: boolean) {
isMentionPickerOpen = v;
},
get mentionQuery() {
return mentionQuery;
},
set mentionQuery(v: string) {
mentionQuery = v;
},
get isWorkingDirectoryPickerOpen() {
return isWorkingDirectoryPickerOpen;
},
set isWorkingDirectoryPickerOpen(v: boolean) {
isWorkingDirectoryPickerOpen = v;
},
get workingDirectoryQuery() {
return workingDirectoryQuery;
},
set workingDirectoryQuery(v: string) {
workingDirectoryQuery = v;
},
get availableCommands() {
return availableCommands;
},
get mentionScopePath() {
return mentionScopePath;
},
handleInput,
// True when a picker consumed the event, so the form skips submit.
handleKeydown,
dispatchCommand,
handleCommandSelect,
handleCommandPickerClose,
handleMentionPickerClose,
handlePromptPickerClose,
handleWorkingDirectoryOpen,
handleWorkingDirectoryClose,
openPromptPicker() {
isPromptPickerOpen = true;
},
closePromptPicker() {
isPromptPickerOpen = false;
promptSearchQuery = '';
}
};
}
export type UseChatFormPickersReturn = ReturnType<typeof useChatFormPickers>;
@@ -0,0 +1,67 @@
import { debounce } from '$lib/utils/debounce';
/**
* Shared debounced async-search machinery for the chat-form pickers:
* AbortController + sequence counter to discard stale responses, a
* debounce, and a live `isSearching` flag.
*/
export interface UseDebouncedSearchOptions {
debounceMs: number;
/** Fire-time guard: a scheduled call that outlives a reset is dropped. */
canRun: () => boolean;
/** Live query, used to drop a scheduled call whose query changed. */
getQuery: () => string;
/** Perform the search and commit results; bail out when `isCurrent()` is false. */
run: (query: string, signal: AbortSignal, isCurrent: () => boolean) => void | Promise<void>;
}
export function useDebouncedSearch(opts: UseDebouncedSearchOptions) {
let controller: AbortController | null = null;
let searchSeq = 0;
let isSearching = $state(false);
function isCurrent(seq: number) {
return seq === searchSeq;
}
function cancel() {
controller?.abort();
searchSeq++;
isSearching = false;
}
const schedule = debounce((query: string) => {
if (!opts.canRun() || query !== opts.getQuery().trim()) return;
void start(query);
}, opts.debounceMs);
async function start(query: string) {
cancel();
const fresh = new AbortController();
controller = fresh;
const mySeq = ++searchSeq;
isSearching = true;
try {
await opts.run(query, fresh.signal, () => isCurrent(mySeq));
} finally {
if (isCurrent(mySeq)) isSearching = false;
}
}
return {
get isSearching() {
return isSearching;
},
/** Bump the loading flag synchronously (e.g. before the debounce fires). */
setLoading(value: boolean) {
isSearching = value;
},
run(query: string) {
schedule(query);
},
cancel
};
}
export type UseDebouncedSearchReturn = ReturnType<typeof useDebouncedSearch>;
@@ -8,6 +8,7 @@ import {
singleModelName
} from '$lib/stores/models.svelte';
import { isRouterMode } from '$lib/stores/server.svelte';
import { CHAT_INPUT_FOCUS_SELECTOR } from '$lib/constants';
import { filterModelOptions, groupModelOptions } from '$lib/components/app/models/utils';
import type { ModelOption } from '$lib/types/models';
@@ -139,11 +140,9 @@ export function useModelsSelector(opts: UseModelsSelectorOptions): UseModelsSele
handleOpenChange(false);
requestAnimationFrame(() => {
const textarea = document.querySelector<HTMLTextAreaElement>(
'[data-slot="chat-form"] textarea'
);
const input = document.querySelector<HTMLElement>(CHAT_INPUT_FOCUS_SELECTOR);
textarea?.focus({ preventScroll: true });
input?.focus({ preventScroll: true });
});
}
@@ -0,0 +1,108 @@
import { KeyboardKey } from '$lib/enums';
/**
* Shared keyboard navigation state for the chat-form pickers: a highlighted
* row, a scroll trigger, and Arrow/Escape/Enter handling.
*/
export interface UsePickerNavigationOptions {
/** Gates all key handling. */
isOpen: () => boolean;
count: () => number;
/**
* Resolve the row to highlight for a movement step, or -1 when no move
* is possible. Defaults to plain wraparound across `count()`.
*/
step?: (from: number, dir: 1 | -1) => number;
onClose: () => void;
/** Called on Enter when `hoveredIndex` points at a selectable row. */
onSelect: (index: number) => void;
}
function wrapStep(from: number, dir: 1 | -1, count: number): number {
return dir === 1 ? (from + 1) % count : from <= 0 ? count - 1 : from - 1;
}
export function usePickerNavigation(opts: UsePickerNavigationOptions) {
let hoveredIndex = $state(-1);
let scrollTrigger = $state(0);
function resolve(from: number, dir: 1 | -1): number {
const n = opts.count();
if (n === 0) return -1;
if (opts.step) return opts.step(from, dir);
return wrapStep(from, dir, n);
}
function move(dir: 1 | -1) {
const next = resolve(hoveredIndex, dir);
if (next >= 0) {
hoveredIndex = next;
scrollTrigger++;
}
}
/** Reset the highlight without bumping the scroll trigger. */
function reset(index: number) {
hoveredIndex = index;
}
/** Bump the scroll trigger without moving the highlight. */
function bumpScroll() {
scrollTrigger++;
}
/** Mouse hover highlights a row but must NOT bump the scroll trigger. */
function setHover(index: number) {
hoveredIndex = index;
}
function handleKeydown(event: KeyboardEvent): boolean {
if (!opts.isOpen()) return false;
if (event.key === KeyboardKey.ESCAPE) {
event.preventDefault();
opts.onClose();
return true;
}
if (event.key === KeyboardKey.ARROW_DOWN) {
event.preventDefault();
move(1);
return true;
}
if (event.key === KeyboardKey.ARROW_UP) {
event.preventDefault();
move(-1);
return true;
}
if (event.key === KeyboardKey.ENTER) {
if (hoveredIndex >= 0 && hoveredIndex < opts.count()) {
event.preventDefault();
opts.onSelect(hoveredIndex);
return true;
}
// No selectable row - let the caller's Enter-to-submit run.
return false;
}
return false;
}
return {
get hoveredIndex() {
return hoveredIndex;
},
get scrollTrigger() {
return scrollTrigger;
},
reset,
setHover,
move,
bumpScroll,
handleKeydown
};
}
export type UsePickerNavigationReturn = ReturnType<typeof usePickerNavigation>;
@@ -0,0 +1,47 @@
import { untrack } from 'svelte';
/**
* Scrolls the highlighted row of a picker list into view when the scroll
* trigger is bumped, without scrolling on mouse hover or result
* replacement.
*/
export interface UseScrollActiveRowOptions {
/** Counter bumped by keyboard nav; `undefined` disables the effect. */
getTrigger: () => number | undefined;
getContainer: () => HTMLDivElement | null;
getIndex: () => number;
getCount: () => number;
/** Attribute prefix, e.g. 'picker' for `[data-picker-index="0"]`. */
dataIndex: string;
}
export function useScrollActiveRow(opts: UseScrollActiveRowOptions) {
let lastTrigger: number | null = null;
$effect(() => {
const trigger = opts.getTrigger();
if (trigger === undefined) return;
// Skip the initial run on mount: the list opens with the first row
// already in view, and scrolling here fires before the popover is
// positioned, which would scroll the whole page to the top.
if (lastTrigger === null) {
lastTrigger = trigger;
return;
}
if (trigger === lastTrigger) return;
lastTrigger = trigger;
untrack(() => {
const container = opts.getContainer();
const index = opts.getIndex();
if (!container || index < 0 || index >= opts.getCount()) return;
const row = container.querySelector(
`[data-${opts.dataIndex}-index="${index}"]`
) as HTMLElement | null;
row?.scrollIntoView({ block: 'nearest', inline: 'nearest' });
});
});
}
export type UseScrollActiveRowReturn = ReturnType<typeof useScrollActiveRow>;
+22 -1
View File
@@ -1,7 +1,12 @@
import { base } from '$app/paths';
import { SvelteMap, SvelteSet } from 'svelte/reactivity';
import { toast } from 'svelte-sonner';
import { ServerModelStatus, ServerModelsSseEventType, ModelModality } from '$lib/enums';
import {
ServerModelStatus,
ServerModelsSseEventType,
ModelModality,
FileTypeCategory
} from '$lib/enums';
import { ModelsService } from '$lib/services/models.service';
import { PropsService } from '$lib/services/props.service';
import { serverStore, isRouterMode } from '$lib/stores/server.svelte';
@@ -394,6 +399,7 @@ class ModelsStore {
model: modelId,
description: details?.description,
capabilities: rawCapabilities.filter((value: unknown): value is string => Boolean(value)),
modalities: this.buildArchitectureModalities(item.architecture),
details: details?.details,
meta: item.meta ?? null,
parsedId: ModelsService.parseModelId(modelId),
@@ -1000,6 +1006,21 @@ class ModelsStore {
};
}
/** Map the router modalities, the only source available while a model is not loaded. */
private buildArchitectureModalities(
architecture: ApiModelDataEntry['architecture']
): ModelModalities | undefined {
if (!architecture) return undefined;
const inputs = architecture.input_modalities;
return {
vision: inputs.includes(FileTypeCategory.IMAGE),
audio: inputs.includes(FileTypeCategory.AUDIO),
video: inputs.includes(FileTypeCategory.VIDEO)
};
}
clear(): void {
this.unsubscribeStatus();
this.statusWaiters.forEach((waiter) => waiter.reject(new Error('Models store cleared')));
@@ -135,6 +135,30 @@ class SettingsStore {
...savedVal
};
// Migrate the legacy render keys into `renderContentAsRawText`
// (inverted semantics: the old keys opted INTO markdown). Any
// explicit raw-text preference wins when the legacy keys disagree.
const LEGACY_MARKDOWN_KEYS = ['renderUserContentAsMarkdown', 'renderThinkingAsMarkdown'];
const LEGACY_RAW_TEXT_KEY = 'renderUserContentAsRawText'; // this branch's intermediate key
const legacyKeys = [...LEGACY_MARKDOWN_KEYS, LEGACY_RAW_TEXT_KEY].filter(
(key) => key in savedVal
);
if (legacyKeys.length > 0) {
if (!(SETTINGS_KEYS.RENDER_CONTENT_AS_RAW_TEXT in savedVal)) {
if (LEGACY_RAW_TEXT_KEY in savedVal) {
this.config[SETTINGS_KEYS.RENDER_CONTENT_AS_RAW_TEXT] = savedVal[LEGACY_RAW_TEXT_KEY];
} else {
this.config[SETTINGS_KEYS.RENDER_CONTENT_AS_RAW_TEXT] = LEGACY_MARKDOWN_KEYS.filter(
(key) => key in savedVal
).some((key) => savedVal[key] === false);
}
}
for (const key of legacyKeys) {
delete (this.config as Record<string, unknown>)[key];
}
this.saveConfig();
}
// Default sendOnEnter to false on mobile when the user has no saved preference
if (!(SETTINGS_KEYS.SEND_ON_ENTER in savedVal)) {
if (isMobile.current) {
+11
View File
@@ -98,10 +98,21 @@ export interface ApiModelDataEntry {
aliases?: string[];
/** Informational tags for this model */
tags?: string[];
/** Modality capabilities, reported by the router for every model regardless of load state */
architecture?: ApiModelArchitecture;
/** Legacy meta field (may be present in older responses) */
meta?: Record<string, unknown> | null;
}
/**
* Modality capabilities of a model, as advertised by the ROUTER /models endpoint.
* Read from the model manifest, so it is available before the model is loaded.
*/
export interface ApiModelArchitecture {
/** Accepted input modalities, always contains "text" */
input_modalities: string[];
}
/**
* Load stage reported by the /models/sse feed, in load order.
*/
+25 -1
View File
@@ -1,4 +1,4 @@
import type { ErrorDialogType } from '$lib/enums';
import type { ChatFormCommandAction, ErrorDialogType, FileMentionEntryType } from '$lib/enums';
import type { ApiChatCompletionToolCall } from './api';
import type { DatabaseMessage, DatabaseMessageExtra } from './database';
@@ -166,3 +166,27 @@ export interface FileProcessingResult {
extras: DatabaseMessageExtra[];
emptyFiles: string[];
}
/**
* A file or folder picked in the @-mention picker. `path` is the absolute
* server-side path; `name` is the basename.
*/
export interface FileMentionEntry {
path: string;
name: string;
type: FileMentionEntryType;
}
/**
* A slash command surfaced by the `/` command picker. `disabled` marks a
* command whose backing capability is unavailable (e.g. `/prompt` when no
* MCP server exposes prompts): visible but greyed out and not selectable.
*/
export interface ChatFormCommand {
name: string;
description: string;
/** Extra search terms that should match this command in the picker. */
keywords?: string[];
action: ChatFormCommandAction;
disabled: boolean;
}
+3 -1
View File
@@ -53,7 +53,9 @@ export type {
LiveProcessingStats,
LiveGenerationStats,
AttachmentDisplayItemsOptions,
FileProcessingResult
FileProcessingResult,
FileMentionEntry,
ChatFormCommand
} from './chat.d';
// Database types
+6
View File
@@ -31,6 +31,9 @@ export interface SettingsEntry {
radioOptions?: Array<{ value: string; label: string; key: string; isExperimental?: boolean }>;
isExperimental?: boolean;
isPositiveInteger?: boolean;
placeholder?: string;
min?: number;
max?: number;
dependsOn?: string;
sync?: {
serverKey: string;
@@ -52,6 +55,9 @@ export interface SettingsFieldConfig {
type: SettingsFieldType;
isExperimental?: boolean;
isPositiveInteger?: boolean;
placeholder?: string;
min?: number;
max?: number;
dependsOn?: string;
help?: string;
options?: Array<{ value: string; label: string; icon?: typeof Icon }>;
+31
View File
@@ -0,0 +1,31 @@
/**
* Slash-command token detection for the chat form. Valid only at offset 0.
*/
export function findCommandToken(
value: string
): { name: string; args: string; end: number } | null {
if (!value.startsWith('/')) return null;
const rest = value.slice(1);
const spaceIdx = rest.search(/\s/);
const name = spaceIdx === -1 ? rest : rest.slice(0, spaceIdx);
const args = spaceIdx === -1 ? '' : rest.slice(spaceIdx + 1);
return { name, args, end: value.length };
}
/**
* Stable signature of a slash-command token for use as a "dismissed"
* marker: while the picker is closed and this exact token is still intact,
* the picker does not re-open on in-token edits.
*/
export interface CommandDismissSnapshot {
name: string;
args: string;
}
export function takeCommandDismissSnapshot(value: string): CommandDismissSnapshot | null {
const token = findCommandToken(value);
if (!token) return null;
return { name: token.name, args: token.args };
}
+151
View File
@@ -0,0 +1,151 @@
/**
* Shared `file_glob_search` runners with a short-lived result cache, so a
* repeated query for the same (type, path, glob, depth) reuses the last
* result instead of re-walking the tree.
*/
import { BuiltInTool, GlobSearchType } from '$lib/enums';
import { ToolsService } from '$lib/services/tools.service';
import {
GLOB_WILDCARD,
PATH_NAV_MAX_DEPTH,
PATH_SEPARATOR,
WINDOWS_SEPARATOR
} from '$lib/constants';
import { lastPathSegment } from './path-display';
import {
buildGlobSearchArgs,
joinPath,
rankEntries,
type GlobEntry,
type GlobSearchArgs
} from './working-directory';
const SEARCH_CACHE_TTL_MS = 2000;
interface CacheEntry {
results: GlobEntry[];
base: string;
at: number;
}
const searchCache = new Map<string, CacheEntry>();
export interface GlobSearchResult {
base: string;
entries: GlobEntry[];
error?: string;
}
export async function runGlobSearch(
args: GlobSearchArgs,
type: GlobSearchType,
limit: number,
signal: AbortSignal
): Promise<GlobSearchResult> {
const key = `${type}\u0000${args.path}\u0000${args.include}\u0000${args.maxDepth}\u0000${limit}`;
const cached = searchCache.get(key);
if (cached && Date.now() - cached.at < SEARCH_CACHE_TTL_MS) {
return { base: cached.base, entries: cached.results };
}
const res = await ToolsService.executeToolRaw(
BuiltInTool.FILE_GLOB_SEARCH,
{ path: args.path, type, include: args.include, max_depth: args.maxDepth, limit },
signal
);
if (typeof res.error === 'string') return { base: '', entries: [], error: res.error };
const base = typeof res.base === 'string' ? res.base : '';
const entries = Array.isArray(res.entries) ? (res.entries as GlobEntry[]) : [];
const now = Date.now();
// prune stale entries so the short-lived cache cannot grow unbounded
for (const [k, v] of searchCache) {
if (now - v.at >= SEARCH_CACHE_TTL_MS) searchCache.delete(k);
}
searchCache.set(key, { results: entries, base, at: now });
return { base, entries };
}
export interface GlobEntryResult {
path: string;
name: string;
type: string;
}
export interface GlobSearchChildOptions {
type?: GlobSearchType;
/** Descend only on a trailing path separator (mention picker); off for
* the WD picker, which descends on any exact match. */
descendOnTrailingSeparator?: boolean;
childMaxDepth?: number;
}
export interface GlobSearchChildResult {
base: string;
args: GlobSearchArgs;
/** Outer ranked entries plus the walked directory's children (absolute). */
entries: GlobEntryResult[];
/** Absolute path of the directory whose children were appended. */
exactDir?: string;
error?: string;
}
function toEntryResult(e: GlobEntry, base: string): GlobEntryResult {
return { path: joinPath(base, e.path), name: lastPathSegment(e.path), type: e.type };
}
/**
* One ranked glob search that may also list the matched directory's
* children, shared by the WD picker (descend on exact match) and the
* mention picker (descend on a trailing `/` or `\`).
*/
export async function runGlobSearchWithChildren(
query: string,
scopePath: string,
searchDepth: number,
limit: number,
signal: AbortSignal,
options: GlobSearchChildOptions = {}
): Promise<GlobSearchChildResult> {
const {
type = GlobSearchType.ALL,
descendOnTrailingSeparator = false,
childMaxDepth = PATH_NAV_MAX_DEPTH
} = options;
const args = buildGlobSearchArgs(query, scopePath, searchDepth);
const res = await runGlobSearch(args, type, limit, signal);
if (res.error) return { base: res.base, args, entries: [], error: res.error };
const ranked = rankEntries(res.entries, args.rankQuery);
const entries = ranked.map((e) => toEntryResult(e, res.base));
const last = args.last;
if (last) {
const wantsDescend = descendOnTrailingSeparator
? query.endsWith(PATH_SEPARATOR) || query.endsWith(WINDOWS_SEPARATOR)
: true;
const exact = ranked.find(
(e) => e.type === 'dir' && lastPathSegment(e.path).toLowerCase() === last.toLowerCase()
);
if (wantsDescend && exact) {
const exactDir = joinPath(res.base, exact.path);
const childRes = await runGlobSearch(
{ path: exactDir, include: GLOB_WILDCARD, maxDepth: childMaxDepth, rankQuery: '' },
type,
limit,
signal
);
if (!childRes.error) {
const children = childRes.entries
.map((e) => toEntryResult(e, childRes.base))
.sort((a, b) => a.path.localeCompare(b.path));
return { base: res.base, args, entries: [...entries, ...children], exactDir };
}
}
}
return { base: res.base, args, entries };
}
+41
View File
@@ -174,13 +174,54 @@ export {
export {
splitPathQuery,
buildCaseInsensitiveGlob,
buildGlobSearchArgs,
rankEntries,
joinPath,
highlightMatch,
type GlobEntry,
type GlobSearchArgs,
type PathQuery
} from './working-directory';
// Shared `file_glob_search` runner with a short-lived result cache
export {
runGlobSearch,
runGlobSearchWithChildren,
type GlobEntryResult,
type GlobSearchResult
} from './glob-search';
// Mention-token detection (for the `@`-triggered file/folder mention picker)
export {
findMentionToken,
takeMentionDismissSnapshot,
type MentionDismissSnapshot
} from './mention-token';
// Slash-command token detection (for the `/`-triggered command picker)
export {
findCommandToken,
takeCommandDismissSnapshot,
type CommandDismissSnapshot
} from './command-token';
// Mention-chip visual contract shared by the rehype file-badge plugin,
// plus the `[name](file://...)` link helpers the mention picker splices in
export {
fileMentionLinkRe,
encodeFileLinkPath,
decodeFileLinkPath,
MENTION_BADGE_CLASSNAME,
MENTION_BADGE_ICON_CLASSNAME,
MENTION_BADGE_SVG_ATTRIBUTES,
MENTION_BADGE_FILE_ICON_PATHS,
MENTION_BADGE_FOLDER_ICON_PATHS,
getMentionBadgeIconPaths,
getMentionBadgeLabel,
buildMentionInsertion,
mentionLinkEndingAt
} from './mention-badge';
// Agentic content utilities (structured section derivation)
export {
deriveAgenticSections,

Some files were not shown because too many files have changed in this diff Show More