feat: add frequency penalty to inference sampling pipeline

- Add FrequencyPenaltyStrategy (logit -= penalty * count)
- Per-task rep_window for penalty history lookup
- Wire through engine, task, executor, API layer
- Add --frequency_penalty and --rep_window to stream_chat.py
- 9 unit tests for frequency penalty strategy
This commit is contained in:
2026-07-17 21:28:31 +08:00
parent a1ea26d367
commit d08a92c7bd
9 changed files with 370 additions and 20 deletions
+30
View File
@@ -75,6 +75,33 @@ class Executor:
temperatures = torch.tensor([t.temperature for t in tasks], device=self.device)
top_ks = torch.tensor([t.top_k for t in tasks], device=self.device)
top_ps = torch.tensor([t.top_p for t in tasks], device=self.device)
freq_penalties = torch.tensor(
[t.frequency_penalty for t in tasks], device=self.device
)
history_lists = []
mask_lists = []
for t in tasks:
window = t.rep_window
prompt_part = t.prompt_ids[-window:]
ids = prompt_part + t.output_ids
history_lists.append(ids)
mask_lists.append([True] * len(ids))
max_len = max(len(h) for h in history_lists)
padded_ids = torch.zeros(
len(tasks), max_len, dtype=torch.long, device=self.device
)
padded_mask = torch.zeros(
len(tasks), max_len, dtype=torch.bool, device=self.device
)
for i, (h, m) in enumerate(zip(history_lists, mask_lists)):
padded_ids[i, : len(h)] = torch.tensor(
h, dtype=torch.long, device=self.device
)
padded_mask[i, : len(m)] = torch.tensor(
m, dtype=torch.bool, device=self.device
)
with torch.inference_mode():
outputs = self.model(
@@ -89,4 +116,7 @@ class Executor:
temperature=temperatures,
top_k=top_ks,
top_p=top_ps,
frequency_penalty=freq_penalties,
input_ids=padded_ids,
input_mask=padded_mask,
).tolist()
+8
View File
@@ -33,6 +33,8 @@ class Task:
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
):
self.task_id = task_id
self.prompt_ids = prompt_ids
@@ -40,6 +42,8 @@ class Task:
self.temperature = temperature
self.top_p = top_p
self.top_k = top_k
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
self.status = TaskStatus.PENDING
self.output_ids: List[int] = []
@@ -92,6 +96,8 @@ class TaskManager:
temperature: float = 1.0,
top_p: float = 1.0,
top_k: int = 50,
frequency_penalty: float = 0.0,
rep_window: int = 64,
stream_callback: Optional[Callable[[str], None]] = None,
) -> str:
task_id = f"task_{int(time.time())}_{uuid.uuid4().hex[:8]}"
@@ -116,6 +122,8 @@ class TaskManager:
temperature=temperature,
top_p=top_p,
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
)
with self._lock: