파일 분리 + STM to LTM

마계닭·2026년 2월 24일

chatbot

목록 보기
3/3

파일 분리


뭔가 많이 추가됐다. 장기기억 + 출력 실험용으로 들어가있는 파일들도 있기에 필요한 것들만 보자.

  • .env: api키가 저장되어있다. 시스템 변수로 저장해두긴했지만, 일단 별도로 저장도 해두었다.

1. config.py

LLM api를 호출하고 기본 세팅을 하는 파일이다.

from google import genai
from google.genai import types
MEM_PATH = "short_term_memory.json"
LONG_PATH = "integrated.json"

TOP_K = 4

STM_MAX = 10
PROMOTE_START = 5
PROMOTE_END = 10
client = genai.Client()
instruction = (
    "You will receive context in this structure:\n"
    "_user query\n<user_input>\n_short_term memory\n<short_term memory>\n"
    "Each section is labeled and appears on its own line. "
    "Use short_term memory for recent turns. Answer the user query based on the provided context."
)

api_model = "gemini-2.5-flash"

setting_config = types.GenerateContentConfig(
    thinking_config=types.ThinkingConfig(thinking_budget=0), # Disables thinking
    system_instruction=instruction,
    temperature=0.9,
    max_output_tokens=200,  
    stop_sequences=["User:", "사용자:", "\n\n"], 
    )

우선 단기기억(MEM_PATH)와 장기기억(LONG_PATH)의 경로, TOP_K를 몇으로 지정할지가 저장되어있다.

STM_MAX와 PROMOTE_START, PROMOTE_END는 이후 설명할 단기기억을 장기기억으로 넘기는 개수 기준이다. 즉, 단기기억이 10개가 쌓일 경우 5번째에서 10번째까지의 기억을 장기기억으로 넘길 것이다.

instructure에선 short_term을 넘겨주는 사실을 명확하게 해두었고, config는 기존과 크게 달라진점은 없다.

2. memory_format.py

def build_contents(turns, user_input=None):
    contents = []
    turns = turns or []

    for t in turns:
        contents.append(f"User: {t['user']}")
        contents.append(f"Assistant: {t['assistant']}")

    if user_input is not None:
        contents.append(f"User: {user_input}")

    return "\n".join(contents)

build_contents를 STM와 LTM 모두에서 사용하기에 별도로 분리해두었다.
추후에 GPT API에게 전송하기 위해서 User과 Assistant를 명시적으로 표시해두었다.

3. short_term.py

import os
import json
from datetime import datetime, timezone


def load_memory(path):
    if not os.path.exists(path):
        return []

    try:
        with open(path, "r", encoding="utf-8") as f:
            data = json.load(f)
        if not isinstance(data, list):
            return []
        turn = []
        for t in data:
            if not isinstance(t, dict):
                continue
            if "user" not in t or "assistant" not in t:
                continue
            created_at = t.get("created_at")
            turn.append({
                "id": len(turn),
                "user": str(t["user"]),
                "assistant": str(t["assistant"]),
                "created_at": created_at,
            })
        return turn
    except (json.JSONDecodeError, OSError):
        return []


def save_memory(path, memory):
    tmp_path = path + ".tmp"
    with open(tmp_path, "w", encoding="utf-8") as f:
        json.dump(memory, f, ensure_ascii=False, indent=2)
    os.replace(tmp_path, path)


def new_turn(user_input, assistant_reply):
    return {
        "id": 0,
        "user": user_input,
        "assistant": assistant_reply,
        "created_at": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
    }

단기기억을 저장하는 파일이다. 대부분 전에 만들었던 부분과 동일하고, created_at이라는 해당 QA가 만들어진 시간을 저장하는 부분을 추가했다.(추후에 사용할 예정)

4. llm_api.py

from google import genai
from google.genai import types
import memory_format
from config import *

def call_gemini(user_input, memory, long_term_turns=None):
    recent_memory = memory[-TOP_K:]
    sections = []

    user_section = "_user query\n" + user_input

    short_term_contents = memory_format.build_contents(recent_memory)
    short_term_section = "_short_term memory\n" + (short_term_contents if short_term_contents else "(empty)")

    sections.append(user_section)
    sections.append(short_term_section)

    user_content = "\n".join(sections)
    
    response = client.models.generate_content(
        model=api_model,
        contents=user_content,
        config=setting_config,
    )

    return response.text.strip()

실제 api를 호출하는 구간이다.

5. main.py

from config import MEM_PATH, LONG_PATH, STM_MAX, PROMOTE_START, PROMOTE_END, TOP_K
import short_term
import llm_api
import long_term
import stm_to_ltm

def main():
    memory = short_term.load_memory(MEM_PATH)

    print("대화 시작 (종료는 Ctrl+C) \n")

    try:
        while True:
            user_input = input("Gemini에게 물어보기: ").strip()
            if not user_input:
                continue
          

            reply = llm_api.call_gemini(user_input, memory)
            print(f"Gemini: {reply}\n")

            new_turn = short_term.new_turn(user_input, reply)
            new_turn["id"] = len(memory)
            memory.append(new_turn)
            short_term.save_memory(MEM_PATH, memory)
            memory = stm_to_ltm.promote_stm_to_ltm(
                MEM_PATH,
                LONG_PATH,
                max_stm=STM_MAX,
                promote_start=PROMOTE_START,
                promote_end=PROMOTE_END,
            )
        
    except KeyboardInterrupt:
        print("\n종료")

if __name__ == "__main__":
    main()

앞의 부분들을 통합 + 다음에 설명될 장기기억으로 변환을 통합해둔 파트이다.

STM to LTM

이제 단기기억이 일정 이상 쌓이면 장기기억으로 변환해주는 파트를 만들어볼 것이다.
3가지 함수로 구현해보았다.

_next_ltm_index

def _next_ltm_index(items):
    if not items:
        return 0
    last = items[-1]
    item_id = last.get("id", "")
    if isinstance(item_id, int):
        return item_id + 1
    if isinstance(item_id, str):
        if item_id.isdigit():
            return int(item_id) + 1
    return len(items)

STM이 아닌 LTM에서 해당 QA의 index를 구하는 파트이다.
현재 LTM에는 항상 오름차순으로 기억이 추가만 된다고 가정해둔 상태이다.(즉, 기억의 중간 삽입이나 삭제가 존재하지 않는다.)
따라서 가장 마지막 item의 id값 + 1로 현재 기억의 id값을 결정한다.

append_to_long_term

def append_to_long_term(integrated_path, turns):
    """
    turns: list of {id, user, assistant, created_at}
    """
    if not os.path.exists(integrated_path):
        data = {"items": []}
    else:
        try:
            with open(integrated_path, "r", encoding="utf-8") as f:
                data = json.load(f)
            if "items" not in data or not isinstance(data["items"], list):
                data = {"items": []}
        except (json.JSONDecodeError, OSError):
            data = {"items": []}

    next_index = _next_ltm_index(data["items"])
    for t in turns:
        record = {
            "id": next_index,
            "user_query": t["user"],
            "answer_query": t["assistant"],
            "created_at": t.get("created_at"),
        }
        data["items"].append(record)
        next_index += 1

    tmp_path = integrated_path + ".tmp"
    with open(tmp_path, "w", encoding="utf-8") as f:
        json.dump(data, f, ensure_ascii=False, indent=2)
    os.replace(tmp_path, integrated_path)

파일에 ltm을 추가하는 파트이다.
integrated_path(LTM 경로)의 파일을 열고, 위의 _next_ltm_index에서 가져온 id와 STM에서 필요한 정보들을 가져와서 그대로 LTM에 넣는다.
데이터 안정성을 위해서 tmp에 저장해두었다가 모두 완료된 이후 원래의 파일에 덮어쓰는 방식을 사용한다.

promote_stm_to_ltm

def promote_stm_to_ltm(stm_path, ltm_path, max_stm=10, promote_start=5):
    turns = short_term.load_memory(stm_path)

    if len(turns) <= max_stm:
        return turns

    excess = len(turns) - max_stm
    promote_count = promote_start + excess + 1
    if promote_count < 0:
        promote_count = 0
    if promote_count > len(turns):
        promote_count = len(turns)

    promote = turns[:promote_count]
    remaining = turns[promote_count:]

    append_to_long_term(ltm_path, promote)

    for i, t in enumerate(remaining):
        t["id"] = i

    short_term.save_memory(stm_path, remaining)
    return remaining

stm쪽을 정리해주는 함수다.
stm의 최대 개수인 max_stm(여기선 10)을 초과한 개수를 구하고, 가장 오래된 0번째 기억부터 (promote_start(여기선 5) + 초과된 개수)까지의 기억을 ltm으로 넘긴다. 이후 남은 stm들의 id를 재정리해준다.

Conclusion

테스트를 해보고싶지만 STM에 어떤 데이터를 넣을지를 고민중이다. 테스트한 이후에 추가할 수 있다면 추가하도록 하겠다.

profile
뉴비

0개의 댓글