File size: 5,455 Bytes
dbc6675
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
"""
Random baseline tokenizer for planning states.

This tokenizer intentionally avoids learned or graph-aware structure.

It is designed as a weak but still task-sensible baseline:
- the state is represented only through coarse predicate-count information
- the goal is represented as a random bag of goal atoms

Compared with a stronger grounded-atom random baseline, this removes most
object-level state structure while still giving the downstream model the
minimum conditioning needed for problem-specific planning.
"""

import hashlib
import json
import logging
import re

import numpy as np

from code.tokenization.base import TokenizationStrategy

logger = logging.getLogger(__name__)

_PREDICATE_REGEX = re.compile(r"\(([\w-]+(?: [\w-]+)*)\)")


def _normalize_atoms(raw_atoms: list[str]) -> list[str]:
    """Normalize atoms by stripping parentheses and sorting lexical tokens."""
    normalized: list[str] = []
    for atom in raw_atoms:
        if not atom:
            continue
        matches = _PREDICATE_REGEX.findall(atom)
        if matches:
            normalized.extend(m.strip().lower() for m in matches if m.strip())
            continue
        clean = atom.replace("(", "").replace(")", "").strip().lower()
        if clean:
            normalized.append(clean)
    return sorted(normalized)


def _predicate_names(normalized_atoms: list[str]) -> list[str]:
    predicates: list[str] = []
    for atom in normalized_atoms:
        parts = atom.split()
        if parts:
            predicates.append(parts[0])
    return sorted(predicates)


class RandomTokenizer(TokenizationStrategy):
    """
    Deterministic random baseline tokenizer.

    The embedding is a deterministic random baseline with asymmetric inputs:

    - state embedding: random bag of state predicate names only
    - goal embedding: random bag of exact goal atoms
    - final vectors are normalized sums

    This is intentionally weaker than a grounded-atom random baseline:
    it keeps only coarse state composition, while the explicit goal vector
    carries the problem-specific conditioning.
    """

    def __init__(self, random_dim: int = 128, seed: int = 42, normalize: bool = True):
        super().__init__(name="Random")
        self.random_dim = int(random_dim)
        self.seed = int(seed)
        self.normalize = bool(normalize)
        self.embedding_dim = self.random_dim

    def fit(
        self,
        domain_pddl_path: str,
        train_states_dir: str,
        train_pddl_dir: str,
    ) -> None:
        """
        No-op fit to match tokenizer interface.

        Args are accepted for API compatibility with other tokenizers.
        """
        _ = (domain_pddl_path, train_states_dir, train_pddl_dir)
        self._is_fitted = True
        logger.info(
            f"[{self.name}] Ready with dim={self.random_dim}, seed={self.seed}, "
            f"normalize={self.normalize}"
        )

    def _vector_for_token(self, token: str) -> np.ndarray:
        digest = hashlib.sha256(token.encode("utf-8")).digest()
        local_seed = int.from_bytes(digest[:8], "big") ^ (self.seed & 0xFFFFFFFFFFFFFFFF)
        rng = np.random.default_rng(local_seed)
        vec = rng.standard_normal(self.random_dim, dtype=np.float32)
        return vec.astype(np.float32)

    def _embed_tokens(self, tokens: list[str]) -> np.ndarray:
        vec = np.zeros(self.random_dim, dtype=np.float32)
        for token in tokens:
            vec += self._vector_for_token(token)

        if self.normalize:
            norm = float(np.linalg.norm(vec))
            if norm > 0:
                vec = vec / norm

        return vec.astype(np.float32)

    def transform_state(
        self,
        state_atoms: list[str],
        goal_atoms: list[str],
        objects: list[str],
    ) -> np.ndarray:
        self._check_fitted()
        state_norm = _normalize_atoms(state_atoms)
        _ = (goal_atoms, objects)

        state_preds = _predicate_names(state_norm)
        tokens = [f"state_pred:{pred}" for pred in state_preds]
        return self._embed_tokens(tokens)

    def transform_goal(
        self,
        goal_atoms: list[str],
        objects: list[str],
    ) -> np.ndarray:
        self._check_fitted()
        goal_norm = _normalize_atoms(goal_atoms)
        _ = objects
        tokens = [f"goal:{atom}" for atom in goal_norm]
        return self._embed_tokens(tokens)

    def get_embedding_dim(self) -> int:
        self._check_fitted()
        return self.random_dim

    def save_vocabulary(self, filepath: str) -> None:
        self._check_fitted()
        payload = {
            "random_dim": self.random_dim,
            "seed": self.seed,
            "normalize": self.normalize,
        }
        with open(filepath, "w") as f:
            json.dump(payload, f, indent=2)
        logger.info(f"[{self.name}] Saved config to {filepath}")

    def load_vocabulary(self, filepath: str) -> None:
        with open(filepath, "r") as f:
            payload = json.load(f)
        self.random_dim = int(payload.get("random_dim", self.random_dim))
        self.seed = int(payload.get("seed", self.seed))
        self.normalize = bool(payload.get("normalize", self.normalize))
        self.embedding_dim = self.random_dim
        self._is_fitted = True
        logger.info(
            f"[{self.name}] Loaded config from {filepath} "
            f"(dim={self.random_dim}, seed={self.seed}, normalize={self.normalize})"
        )