怎麼又要估算呢?不是 Day 26 就完成了嗎?
先前,我把 estimate_token() 直接放在 agent.py 裡,而現在,我打算建立一個獨立管理上下文的資料夾 context,我決定把它移出來。
當然不移F也是無所謂的,只是為了提升專案的結構性。
然後,從 Day 25 就說隔天要來做的上下文壓縮總算是開始了,今天先處理它的第一個部分!
直接把整個函數以及正則從 agent.py 剪下貼到新的 token.py 裡:
# src/meowgent/context/token.py
import re
from typing import Optional
RX_CJK_PATTERN = re.compile(r'[\u4e00-\u9fff\u3000-\u303f\uff00-\uffef]')
def estimate_token(text: str) -> int:
""" 估算出傳入文字的 token 量 """
# 中日韓文、全形標點
cjk_count = len(RX_CJK_PATTERN.findall(text))
# 其餘字元(英文、數字、程式碼標點、半形空白等)
other_chars_count = len(text) - cjk_count
# 中文每字約 1.3 token,其餘每 3.5 字元約 1 token
return int((cjk_count * 1.3) + (other_chars_count / 3.5))
然後還有一個我之前沒注意到的部分,「圖片」的 token 數沒有被估算到(但 ollama 回傳的準確 token 還是有包括的)。
所以,我們要新增一個函數來做計算,
這部分可以做得很簡單也可以做得很複雜,這邊先來看一下「圖片是如何轉為 token」(講一下到底有多複雜,所以我偷懶用最簡單的方式):
圖像尺寸為 $W×H$,分別對應「寬」、「高」。
也就是說,要算出這個 N 就可以得知圖片 token,所以,我們要得知圖片的解析度、以及 P(通常介於 14 到 16),而我們還需要得知它有沒有進行壓縮(否則數據會差好幾倍)。
假設,我們真的取得了這些數據,估算出了圖片較準確的 token 佔用,但這個估算值運用的地方只在「貼上圖片到 ollama 給出回傳前作用」,接著就能得到 ollama 回傳的 token 了,可以說是非常不划算。
所以,說了這麼多,我們直接「假設每張傳入的圖就是 800 個 token」:
這邊把這個函數一樣放在 token.py 裡,讓呼叫者傳送圖片串列進來(input_prompt.py 中的 attached_images),得到長度把它乘上 800。
# src/meowgent/context/token.py
...
def images_token(images: Optional[list[str]] = None):
""" 估算出圖片的 token 量 """
if images:
return 800 * len(images)
else:
return 0
在 get_context_status_text() 加入 images 參數,然後呼叫 images_token:
# src/meowgent/agent.py
class Agent():
...
def get_context_status_text(..., images: Optional[list[str]] = None) -> ...:
...
if self.true_token:
...
input_token = (
estimate_token(user_input) if user_input else 0
) + images_token(images)
else:
...
total_token = estimate_token(total_text) + images_token(images)
在調用 get_context_status_text() 處傳入參數:
這邊剛好要傳入的
attached_images就在input_prompt.py中。
# src/meowgent/cli/input_prompt.py
def get_input(...) -> ...:
...
def _dynamic_rprompt():
return rprompt(..., attached_images)
這麼一來上下文用量的部分就正式完成了,要正式進到「壓縮」的部分了!
先看一下檔案結構:
src/meowgent/
├── ...
└── context
├── __init__.py
├── token.py # 估算 token
├── pruner.py # 第一層壓縮
└── compactor.py # 第二層壓縮
# src/meowgent/context/__init__.py from .token import estimate_token, images_token from .pruner import pruner_tool_text
壓縮的部分,分為兩層,pruner.py - 「裁切」和 compactor.py - 「語意滾動壓縮」。
第一層的 pruner.py 專注於將「回傳量過大的工具調用結果切掉」,且會設定一個限制,只讓系統裁切 n 輪以前的(確保剛讀取到的完整程式碼能順利支援當前輪次的修改與追問)。
這是一個「快速」、且「不耗費模型算力」的簡易的壓縮方式,我設定永遠執行此壓縮方式。
那具體是如何裁切呢?
不能所有情況都無腦做裁切,兩種情況:
- 太多行。
- 太多字。
字數來做限制很好理解,那行數呢?
主要是針對如列表式的工具(如list_file)或讀取程式碼的部分。
但是,這只針對「工具調用的回傳」,使用者的問題、模型的回答還是會持續佔用上下文,這時,就需要第二層發力了!compactor.py 是用模型來做「統整」,用 prompt 指定模型生成特定、統一的格式來讓模型了解前面發生的事情。
但這麼做相當於多讓模型做一次回答,相當耗時,所以只在「上下文窗口快滿時才執行」。
除了這兩種基礎的壓縮方式,還有更進階的,例如讓模型編寫「長期記憶文件」,讓模型在回答之前閱讀。
篇幅關係本系列就不實作了。
剛才提到裁切只會對 n 輪以前做裁切,而一輪指的是「從使用者給出問題」到「模型給出答覆(或不回覆)」。
我們要把一輪一輪給區隔出來,可以透過記錄使用者訊息在 history_messages 的索引來達成。
先建立函數,接收 history_messages,回傳所有使用者輸入的索引列表(先放在 turn_indices):
# src/meowgent/context/pruner.py
from typing import List
import re
def turn_separation(history_messages: List[dict]) -> List[int]:
""" 得到為使用者輸入的索引 """
turn_indices = []
用 for 搭配 enumerate() 做計數,
先判斷 role 是否為 user,但這裡有個陷阱,「不是所有為 user 的都是使用者輸入」,有的是工具調用回傳。
所以,還需要排除掉開頭為 <tool_response> 的工具調用回傳。
回顧一下,工具調用結果長這樣:
<tool_response> [工具 xxx 執行結果]: ... </tool_response>
如果確認真的是使用者輸入,把索引放入 turn_indices。
# src/meowgent/context/pruner.py
...
def turn_separation(...) -> ...:
...
for idx, msg in enumerate(history_messages):
if msg["role"] == "user":
# 排除工具執行的回傳標籤
if not msg["content"].startswith("<tool_response>"):
turn_indices.append(idx)
return turn_indices
確定好第幾輪的問題,接著就可以開始做裁剪的函數了。
接收 history_messages 以及要保留 keep_turns 輪不修改,
然後,調用 _turn_separation() 來獲得使用者輸入索引。
而如果得到的使用者輸入量沒有超過 keep_turns 就不需要繼續,直接回傳原始 history_messages。
# src/meowgent/context/pruner.py
...
def pruner_tool_text(history_messages: List[dict], keep_turns: int = 2):
"""
把 keep_turns 輪以前的工具調用做裁剪
針對:
1. 大於 6 行 -> 留前三後二
2. 大於 300 字 -> 留前後 50
"""
turn_indices = _turn_separation(history_messages)
if len(turn_indices) <= keep_turns: # 還不需要裁剪
return history_messages
然後要來界定要「檢查到哪裡」,用 turn_indices[-keep_turns] 取,然後用 range() 做遍歷:
舉例來說有索引 0 ~ 10,其中 0、3、5、8、10 是使用者輸入(這就是
turn_indices),而剩下的就是需要被檢查的部分。
用turn_indices[-2]就能取出倒數第二個 - 8,然後用range(8)就能得到索引 0 ~ 7,巧妙地避開了 8。
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
cutoff_idx = turn_indices[-keep_turns] # 界定要裁切分界點
# 在 cutoff_idx -1 之前都要檢查
for idx in range(cutoff_idx):
這邊,要先判斷是否為工具執行的回傳,然後檢查是否已經有 ... [省略前期工具 在內(這是等等要放在裁切後的提示),有的話直接跳過此索引。
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
for idx in range(cutoff_idx):
msg = history_messages[idx]
content = msg["content"]
if msg["role"] == "user" and content.startswith("<tool_response>"):
if "... [省略前期工具" in content:
continue
然後又是要用到正則了,目的是把「工具名稱」以及「工具回傳」抽出放在 group(1) 和 group(2),先定義語法:
# src/meowgent/context/pruner.py
from ...
RX_TOOL_RESPONSE = re.compile(
r"<tool_response>\s*\[工具\s*(\w+)\s*執行結果\]:\n([\s\S]*?)\n</tool_response>"
) # 預編譯
用 search() 去尋找,如果沒找到則跳出此索引,找到則取出 group(1) 和 group(2):
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
for ...:
...
if msg["role"] ...:
...
match = RX_TOOL_RESPONSE.search(content)
if not match:
continue
tool_name = match.group(1)
raw_result = match.group(2) # 工具輸出的字串
先做行數的檢查,用 splitlines() 做成分行的串列,用 len() 來判斷行數。
如果超過,用 "\n".join() 來拼出前三行和後兩行,然後拼接出 result:
這邊我發現了寫「分行字串」的另一個方法,之前在 Day 9 時我們用
textwrap.dedent()來做分行效果,但其實用下面的方式也可以有一樣的效果:
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
for ...:
...
if msg["role"] ...:
...
# 檢查行數 -> 超過,留前三後二
lines = raw_result.splitlines()
if len(lines) > 6:
head = "\n".join(lines[:3]) # 保留索引 0~2
tail = "\n".join(lines[-2:]) # 倒數第二個開始留
result = (
f"{head}\n"
f"... [省略前期工具 {tool_name} 的部分輸出]\n"
f"{tail}"
)
再來做字數的檢查,保留前後 50,加上提示:
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
for ...:
...
if msg["role"] ...:
...
if len(lines) > 6:
...
# 檢查字數 -> 超過,留前後 50
elif len(raw_result) > 300:
head = raw_result[:50]
tail = raw_result[-50:]
result = head + f" ... [省略前期工具 {tool_name} 的部分輸出] " + tail
若是 else(未超過行數與字數限制)則直接 continue 跳過此項目;若有裁切則重新拼出工具格式,最後回傳:
# src/meowgent/context/pruner.py
...
def pruner_tool_text(...) -> ...:
...
for ...:
...
if msg["role"] ...:
...
if len(lines) > 6:
...
elif len(raw_result) > 300:
...
else:
continue
msg["content"] = f"<tool_response>\n[工具 {tool_name} 執行結果]:\n{result}\n</tool_response>"
return history_messages
因為是要對多輪紀錄做修改,而它被儲存在 agent.py 的 self.history_messages。
這邊不僅是要呼叫 pruner_tool_text() 寫入 self.history_messages,
還要更新 ollama 回傳出來的準確 token,
首先記錄下原先整個 self.history_messages 佔用的 token(用 estimate_token() 估),然後呼叫 pruner_tool_text() 寫入 self.history_messages,
然後將 token_before 減去估算出的新 token,最後判斷是否需要更新 self.true_token。
# src/meowgent/agent.py
...
class Agent():
def chat(...) -> ...:
...
while turns < self.max_turns:
...
token_before = estimate_token("".join(
m["content"] for m in self.history_messages
))
self.history_messages = pruner_tool_text(self.history_messages)
token_gap = token_before - estimate_token(
"".join(m["content"] for m in self.history_messages)
)
if token_gap > 0 and self.true_token:
self.true_token -= token_gap
上下文壓縮部分今天完成了第一層的裁剪壓縮,明天,來處理比較複雜的第二層 -「語意滾動壓縮」!
(今天總算不是壓線上傳啦,給自己拍拍手!還有兩天,加油加油)