Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 16 additions & 9 deletions memrl/service/value_driven.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from datetime import datetime
from math import isfinite
import random

from memos.mem_os.main import MOS
Expand Down Expand Up @@ -191,20 +192,26 @@ def update(self, memory_id: str, reward: float, next_max_q: Optional[float] = No

old_meta = _meta_to_dict(getattr(item, "metadata", None))
old_q = float(old_meta.get("q_value", self.cfg.q_init_pos))
target = float(reward) + (self.cfg.gamma * float(next_max_q or 0.0))
new_q = (1.0 - self.cfg.alpha) * old_q + self.cfg.alpha * target
reward = float(reward)
next_q = 0.0 if next_max_q is None else float(next_max_q)
alpha, gamma = float(self.cfg.alpha), float(self.cfg.gamma)
old_ma = float(old_meta.get("reward_ma", 0.0))
if not all(isfinite(value) for value in (old_q, reward, next_q, alpha, gamma, old_ma)):
raise ValueError("Q update inputs must be finite")
target = reward + gamma * next_q
new_q = (1.0 - alpha) * old_q + alpha * target
reward_ma = (1.0 - alpha) * old_ma + alpha * reward
if not all(isfinite(value) for value in (target, new_q, reward_ma)):
raise ValueError("Q update arithmetic must remain finite")

# Optional Q floor: prevents Q from dropping below configured minimum.
if getattr(self.cfg, "q_floor", None) is not None:
try:
new_q = max(float(self.cfg.q_floor), float(new_q))
except Exception:
pass
floor = float(self.cfg.q_floor)
if not isfinite(floor):
raise ValueError("Q floor must be finite")
new_q = max(floor, new_q)

visits = int(old_meta.get("q_visits", 0)) + 1
# simple EMA for reward
old_ma = float(old_meta.get("reward_ma", 0.0))
reward_ma = (1.0 - self.cfg.alpha) * old_ma + self.cfg.alpha * float(reward)

new_meta = old_meta | {
"q_value": float(new_q),
Expand Down
86 changes: 86 additions & 0 deletions tests/test_q_value_finite.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
"""Exercise the real updater with an in-memory persistence boundary; no MemOS service."""
import importlib.util
from pathlib import Path
import sys
from types import ModuleType, SimpleNamespace
import unittest
from unittest.mock import patch


def load_module():
stub = ModuleType("memos.mem_os.main")
stub.MOS = object
spec = importlib.util.spec_from_file_location(
"value_driven_under_test",
Path(__file__).resolve().parents[1] / "memrl/service/value_driven.py",
)
module = importlib.util.module_from_spec(spec)
with patch.dict(sys.modules, {"memos.mem_os.main": stub, spec.name: module}):
spec.loader.exec_module(module)
return module


vd = load_module()


class Store:
def __init__(self, metadata):
self.metadata = metadata.copy()
self.writes = []

def get(self, memory_id):
return SimpleNamespace(memory="synthetic", metadata=self.metadata)

def update(self, memory_id, value):
self.writes.append(value)
self.metadata = value["metadata"]


class FiniteUpdateTests(unittest.TestCase):
def make_updater(self, metadata=None, **config):
store = Store(metadata or {"q_value": 0.0})
updater = vd.QValueUpdater(None, "test", vd.RLConfig(**config))
updater._get_text_mem = lambda: store
return updater, store

def test_nonfinite_inputs_never_write(self):
for value in (float("nan"), float("inf"), -float("inf")):
for field in ("reward", "next_max_q", "alpha", "gamma", "q_floor", "q_value", "reward_ma"):
with self.subTest(value=value, field=field):
metadata = {"q_value": 0.0}
config, args = {}, {"reward": 1.0}
if field in ("q_value", "reward_ma"):
metadata[field] = value
elif field in ("reward", "next_max_q"):
args[field] = value
else:
config[field] = value
updater, store = self.make_updater(metadata, **config)
with self.assertRaises(ValueError):
updater.update("m", **args)
self.assertEqual(store.writes, [])

def test_finite_overflow_never_writes(self):
cases = [
({"gamma": 2.0}, {"q_value": 0.0}, 1e308, 1e308),
({"alpha": 2.0}, {"q_value": 0.0}, 1e308, None),
({"alpha": 2.0}, {"q_value": 0.0, "reward_ma": -1e308}, 1e308, None),
]
for config, metadata, reward, next_q in cases:
with self.subTest(config=config, metadata=metadata):
updater, store = self.make_updater(metadata, **config)
with self.assertRaises(ValueError):
updater.update("m", reward, next_q)
self.assertEqual(store.writes, [])

def test_valid_repeated_updates_and_floor(self):
updater, store = self.make_updater(alpha=0.5, gamma=0.5, q_floor=-1.0)
self.assertEqual(updater.update("m", -2.0), -1.0)
self.assertEqual(updater.update("m", 2.0, 2.0), 1.0)
self.assertEqual(store.metadata["q_visits"], 2)
self.assertEqual(store.metadata["reward_ma"], 0.5)
self.assertEqual(store.metadata["last_reward"], 2.0)


if __name__ == "__main__":
unittest.main()