refactor: 统一采样路径为 Strategy + batch tensor,删除 apply_sampling_strategies

- TemperatureStrategy / TopKStrategy / TopPStrategy 支持 Union[float, Tensor]
- SamplingPipeline.sample() 一条调用完成 apply + softmax + multinomial
- 新增 sample() 独立函数作为 scheduler 入口
- scheduler decode 改为 batch tensor 参数传递,支持任意 batch size
- 删除 apply_sampling_strategies(被 sample() 取代)
This commit is contained in:
2026-05-08 19:07:14 +08:00
parent 78dc2bd41c
commit 7ddebf2cd9
3 changed files with 107 additions and 58 deletions
+9 -9
View File
@@ -16,7 +16,7 @@ import torch
from torch import Tensor
from astrai.inference.cache import _STOP, PrefixCacheManager, SlotAllocator
from astrai.inference.sampling import apply_sampling_strategies
from astrai.inference.sampling import sample
from astrai.model.automodel import AutoModel
from astrai.tokenize import AutoTokenizer
@@ -483,14 +483,14 @@ class InferenceScheduler:
)
logits = outputs["logits"][:, -1, :]
next_tokens = []
for i, t in enumerate(tasks):
logit = apply_sampling_strategies(
logits[i : i + 1], t.temperature, t.top_k, t.top_p
)
prob = torch.softmax(logit, dim=-1)
ntok = torch.multinomial(prob, num_samples=1).item()
next_tokens.append(ntok)
next_tokens = sample(
logits,
temperature=torch.tensor(
[t.temperature for t in tasks], device=logits.device
),
top_k=torch.tensor([t.top_k for t in tasks], device=logits.device),
top_p=torch.tensor([t.top_p for t in tasks], device=logits.device),
).tolist()
for t, ntok in zip(tasks, next_tokens):
t.output_ids.append(ntok)