CoolFace
Apppublic

dhruv-punia-bits/memory-compaction-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.py224 linesDownload Raw Back to server
1from __future__ import annotations2 3import uuid4from typing import Any5 6from server.grader import score_step, score_terminal7from server.models import EpisodeState, MemoryItem, MemoryStatus, Observation, Operation, ResetRequest, StepAction, StepResponse, TaskDefinition8from server.tasks import build_episode_payload, build_task_definitions9 10 11class MemoryCompactionEnvironment:12    def __init__(self) -> None:13        self._current_state: EpisodeState | None = None14 15    def list_tasks(self) -> list[TaskDefinition]:16        return build_task_definitions()17 18    def reset(self, request: ResetRequest) -> Observation:19        payload = build_episode_payload(request.difficulty, request.seed)20        self._current_state = EpisodeState(21            episode_id=f"{request.difficulty.value}-{payload['seed']}-{uuid.uuid4().hex[:8]}",22            difficulty=request.difficulty,23            seed=int(payload["seed"]),24            phase="ingest",25            current_index=0,26            token_budget=int(payload["token_budget"]),27            step_count=0,28            recent_turn_window=int(payload["recent_turn_window"]),29            turns=list(payload["turns"]),30            future_queries=list(payload["future_queries"]),31            hidden_gold_memories=list(payload["gold_memories"]),32            hidden_new_memory_ids=[],33            hidden_ambiguous_turn_ids=list(payload.get("ambiguous_turn_ids", [])),34            hidden_low_value_turn_ids=list(payload.get("low_value_turn_ids", [])),35        )36        self._populate_hidden_new_memory_ids()37        return self._build_observation()38 39    def state(self) -> EpisodeState:40        if self._current_state is None:41            raise ValueError("No active episode. Call reset first.")42        return self._current_state43 44    def step(self, action: StepAction) -> StepResponse:45        state = self.state()46        if state.done:47            terminal = score_terminal(state)48            state.latest_reward_breakdown = terminal49            return StepResponse(50                observation=self._build_observation(),51                reward=terminal.total_score,52                done=True,53                reward_breakdown=terminal,54                metadata={"phase": state.phase},55            )56 57        self._apply_action(action)58        state.step_count += 159        step_breakdown = score_step(state, action)60        reward = step_breakdown.total_score61 62        if state.phase == "ingest":63            state.current_index += 164            if state.current_index >= len(state.turns):65                state.phase = "evaluation"66            self._populate_hidden_new_memory_ids()67 68        if state.phase == "evaluation":69            terminal = score_terminal(state)70            reward = round((reward * 0.35) + (terminal.total_score * 0.65), 6)71            step_breakdown.terminal_score = terminal.total_score72            step_breakdown.total_score = reward73            state.done = True74 75        state.latest_reward_breakdown = step_breakdown76        final_reward = round(_strict_reward(reward), 6)77        step_breakdown.total_score = final_reward78        return StepResponse(79            observation=self._build_observation(),80            reward=final_reward,81            done=state.done,82            reward_breakdown=step_breakdown,83            metadata={84                "phase": state.phase,85                "token_usage": state.token_usage(),86                "token_budget": state.token_budget,87            },88        )89 90    def _apply_action(self, action: StepAction) -> None:91        state = self.state()92        if action.summary_text is not None:93            state.working_summary = self._truncate_summary(action.summary_text)94 95        if action.operation == Operation.APPEND_MEMORY:96            for item in action.memory_items:97                self._upsert_memory(item, preserve_history=True)98        elif action.operation == Operation.UPDATE_MEMORY:99            for incoming in action.memory_items:100                self._upsert_memory(incoming, preserve_history=False)101        elif action.operation == Operation.DELETE_MEMORY:102            ids_to_delete = {item.memory_id for item in action.memory_items}103            state.durable_memory = [item for item in state.durable_memory if item.memory_id not in ids_to_delete]104        elif action.operation == Operation.REPLACE_SUMMARY:105            state.working_summary = self._truncate_summary(action.summary_text or "")106 107    def _build_observation(self) -> Observation:108        state = self.state()109        current_turn = state.turns[state.current_index] if state.phase == "ingest" and state.current_index < len(state.turns) else None110        start = max(0, state.current_index - state.recent_turn_window + 1)111        recent_turns = state.turns[start : state.current_index + 1]112        return Observation(113            episode_id=state.episode_id,114            difficulty=state.difficulty,115            phase=state.phase,116            current_turn=current_turn,117            recent_turns=recent_turns,118            working_summary=state.working_summary,119            durable_memory=state.durable_memory,120            token_budget_remaining=state.token_budget_remaining(),121            step_count=state.step_count,122            future_query_queue_size=len(state.future_queries),123            done=state.done,124            info={125                "future_queries": [query.prompt for query in state.future_queries] if state.phase == "evaluation" else [],126                "token_budget": state.token_budget,127                "memory_write_policy": [128                    "Store only durable information likely to matter later.",129                    "Prefer uncertain or summary-only handling for ambiguous statements.",130                    "When a value changes, keep the latest value active and mark the old one superseded.",131                ],132                "ambiguous_turn_ids": state.hidden_ambiguous_turn_ids,133                "low_value_turn_ids": state.hidden_low_value_turn_ids,134            },135        )136 137    def _populate_hidden_new_memory_ids(self) -> None:138        state = self.state()139        payload = build_episode_payload(state.difficulty, state.seed)140        mapping: dict[int, list[str]] = payload["new_memory_ids_by_turn"]  # type: ignore[assignment]141        if state.phase != "ingest" or state.current_index >= len(state.turns):142            state.hidden_new_memory_ids = []143            return144        turn_id = state.turns[state.current_index].turn_id145        state.hidden_new_memory_ids = list(mapping.get(turn_id, []))146 147    def _upsert_memory(self, item: MemoryItem, preserve_history: bool) -> None:148        state = self.state()149        incoming = item.model_copy(deep=True)150        incoming.source_turn_ids = sorted(set(incoming.source_turn_ids))151 152        conflicting_indices: list[int] = []153        for index, existing in enumerate(state.durable_memory):154            if existing.memory_id == incoming.memory_id:155                state.durable_memory[index] = self._merge_memory(existing, incoming)156                self._supersede_relation_conflicts(state.durable_memory[index])157                return158            if (159                existing.status == MemoryStatus.ACTIVE160                and existing.relation_key() == incoming.relation_key()161                and existing.object.lower() != incoming.object.lower()162            ):163                conflicting_indices.append(index)164 165        if conflicting_indices and not incoming.requires_confirmation:166            for index in conflicting_indices:167                conflicting = state.durable_memory[index]168                conflicting.status = MemoryStatus.SUPERSEDED169                conflicting.updated_from_memory_id = conflicting.updated_from_memory_id or incoming.memory_id170 171        if preserve_history or not any(memory.memory_id == incoming.memory_id for memory in state.durable_memory):172            state.durable_memory.append(incoming)173        else:174            for index, existing in enumerate(state.durable_memory):175                if existing.memory_id == incoming.memory_id:176                    state.durable_memory[index] = self._merge_memory(existing, incoming)177                    break178 179        self._supersede_relation_conflicts(incoming)180 181    def _merge_memory(self, existing: MemoryItem, incoming: MemoryItem) -> MemoryItem:182        merged = existing.model_copy(deep=True)183        merged.type = incoming.type184        merged.subject = incoming.subject185        merged.predicate = incoming.predicate186        merged.object = incoming.object187        merged.confidence = incoming.confidence188        merged.source_turn_ids = sorted(set(existing.source_turn_ids + incoming.source_turn_ids))189        merged.source_text = incoming.source_text or existing.source_text190        merged.status = incoming.status191        merged.updated_from_memory_id = incoming.updated_from_memory_id or existing.updated_from_memory_id192        merged.expires_at = incoming.expires_at or existing.expires_at193        merged.importance = incoming.importance194        merged.task_relevance = incoming.task_relevance195        merged.requires_confirmation = incoming.requires_confirmation196        return merged197 198    def _supersede_relation_conflicts(self, incoming: MemoryItem) -> None:199        state = self.state()200        if incoming.status != MemoryStatus.ACTIVE or incoming.requires_confirmation:201            return202        for existing in state.durable_memory:203            if existing.memory_id == incoming.memory_id:204                continue205            if (206                existing.status == MemoryStatus.ACTIVE207                and existing.relation_key() == incoming.relation_key()208                and existing.object.lower() != incoming.object.lower()209            ):210                existing.status = MemoryStatus.SUPERSEDED211                existing.updated_from_memory_id = existing.updated_from_memory_id or incoming.memory_id212 213    def _truncate_summary(self, text: str, max_words: int = 40) -> str:214        return " ".join(text.split()[:max_words])215 216 217def serialise_state(state: EpisodeState) -> dict[str, Any]:218    return state.model_dump()219 220 221def _strict_reward(value: float) -> float:222    epsilon = 0.001223    return max(epsilon, min(1.0 - epsilon, value))224