agent-learning 0.4.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (71) hide show
  1. agent_learning/__init__.py +160 -0
  2. agent_learning/_version.py +3 -0
  3. agent_learning/capture.py +271 -0
  4. agent_learning/classifiers/__init__.py +37 -0
  5. agent_learning/classifiers/base.py +185 -0
  6. agent_learning/classifiers/router.py +236 -0
  7. agent_learning/classifiers/scorers/__init__.py +34 -0
  8. agent_learning/classifiers/scorers/_base.py +162 -0
  9. agent_learning/classifiers/scorers/adherence.py +26 -0
  10. agent_learning/classifiers/scorers/completion.py +26 -0
  11. agent_learning/classifiers/scorers/intent.py +25 -0
  12. agent_learning/cli.py +386 -0
  13. agent_learning/config.py +524 -0
  14. agent_learning/learners/__init__.py +6 -0
  15. agent_learning/learners/base.py +38 -0
  16. agent_learning/learners/reinforce.py +153 -0
  17. agent_learning/metrics/__init__.py +22 -0
  18. agent_learning/metrics/base.py +233 -0
  19. agent_learning/metrics/intent_resolution.py +50 -0
  20. agent_learning/metrics/registry.py +42 -0
  21. agent_learning/metrics/task_adherence.py +42 -0
  22. agent_learning/metrics/task_completion.py +54 -0
  23. agent_learning/policy/__init__.py +7 -0
  24. agent_learning/policy/base.py +59 -0
  25. agent_learning/policy/contextual_softmax.py +243 -0
  26. agent_learning/policy/softmax_bandit.py +157 -0
  27. agent_learning/py.typed +1 -0
  28. agent_learning/rewards/__init__.py +6 -0
  29. agent_learning/rewards/shaping.py +121 -0
  30. agent_learning/rewards/writer.py +130 -0
  31. agent_learning/scorers/__init__.py +187 -0
  32. agent_learning/scorers/base.py +49 -0
  33. agent_learning/scorers/llm/__init__.py +19 -0
  34. agent_learning/scorers/llm/_base.py +126 -0
  35. agent_learning/scorers/llm/adherence.py +21 -0
  36. agent_learning/scorers/llm/completion.py +21 -0
  37. agent_learning/scorers/llm/intent.py +21 -0
  38. agent_learning/scorers/nlp/__init__.py +16 -0
  39. agent_learning/scorers/nlp/_base.py +99 -0
  40. agent_learning/scorers/nlp/adherence.py +17 -0
  41. agent_learning/scorers/nlp/completion.py +17 -0
  42. agent_learning/scorers/nlp/intent.py +17 -0
  43. agent_learning/scorers/nlp_text/__init__.py +26 -0
  44. agent_learning/scorers/nlp_text/_base.py +234 -0
  45. agent_learning/scorers/nlp_text/adherence.py +94 -0
  46. agent_learning/scorers/nlp_text/completion.py +91 -0
  47. agent_learning/scorers/nlp_text/intent.py +59 -0
  48. agent_learning/scorers/slm/__init__.py +25 -0
  49. agent_learning/scorers/slm/_base.py +292 -0
  50. agent_learning/scorers/slm/adherence.py +98 -0
  51. agent_learning/scorers/slm/completion.py +111 -0
  52. agent_learning/scorers/slm/intent.py +80 -0
  53. agent_learning/scorers/stdlib/__init__.py +38 -0
  54. agent_learning/scorers/stdlib/_text.py +87 -0
  55. agent_learning/scorers/stdlib/adherence.py +157 -0
  56. agent_learning/scorers/stdlib/completion.py +117 -0
  57. agent_learning/scorers/stdlib/intent.py +182 -0
  58. agent_learning/storage/__init__.py +14 -0
  59. agent_learning/storage/base.py +156 -0
  60. agent_learning/storage/cosmos.py +506 -0
  61. agent_learning/storage/local.py +353 -0
  62. agent_learning/storage/memory.py +209 -0
  63. agent_learning/training/__init__.py +5 -0
  64. agent_learning/training/runner.py +172 -0
  65. agent_learning/types.py +507 -0
  66. agent_learning-0.4.1.dist-info/METADATA +82 -0
  67. agent_learning-0.4.1.dist-info/RECORD +71 -0
  68. agent_learning-0.4.1.dist-info/WHEEL +5 -0
  69. agent_learning-0.4.1.dist-info/entry_points.txt +2 -0
  70. agent_learning-0.4.1.dist-info/licenses/LICENSE +21 -0
  71. agent_learning-0.4.1.dist-info/top_level.txt +1 -0
agent_learning/cli.py ADDED
@@ -0,0 +1,386 @@
1
+ """Command-line interface for the agent-learning SDK."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import logging
8
+ import sys
9
+ import uuid
10
+ from collections import Counter
11
+ from typing import Any
12
+
13
+ from .policy.softmax_bandit import SoftmaxPolicy
14
+ from .storage.cosmos import get_default_store
15
+ from .training.runner import LearningRunner
16
+ from .types import Action, Episode, MetricName, PolicySnapshot, RewardSource
17
+
18
+ logger = logging.getLogger(__name__)
19
+
20
+ _MAX_EPISODES = 500
21
+
22
+
23
+ def _episode_limit(value: str) -> int:
24
+ limit = int(value)
25
+ if limit < 1 or limit > _MAX_EPISODES:
26
+ raise argparse.ArgumentTypeError(f"limit must be between 1 and {_MAX_EPISODES}")
27
+ return limit
28
+
29
+
30
+ def _build_arg_parser() -> argparse.ArgumentParser:
31
+ parser = argparse.ArgumentParser(prog="agent-learn", description="Native RL CLI for AI agents.")
32
+ sub = parser.add_subparsers(dest="command", required=True)
33
+
34
+ sub.add_parser("list", help="List discovered agent ids and names.")
35
+
36
+ tasks = sub.add_parser("tasks-list", help="List tasks for an agent.")
37
+ tasks.add_argument("agent_id")
38
+
39
+ count = sub.add_parser(
40
+ "task-episodes-count",
41
+ help="Count full learning episodes for an agent.",
42
+ )
43
+ count.add_argument("agent_id")
44
+ count.add_argument("--task-id")
45
+
46
+ episodes = sub.add_parser(
47
+ "task-episodes-list",
48
+ help="Print episodes with score and reward details.",
49
+ )
50
+ episodes.add_argument("agent_id")
51
+ episodes.add_argument("--task-id")
52
+ episodes.add_argument("--limit", type=_episode_limit, default=_MAX_EPISODES)
53
+ episodes.add_argument("--include-incomplete", action="store_true")
54
+
55
+ train = sub.add_parser("train", help="Run one offline learning batch.")
56
+ train.add_argument("--agent-id", required=True)
57
+ train.add_argument("--task-id")
58
+ train.add_argument("--limit", type=_episode_limit, default=200)
59
+ train.add_argument("--start-date")
60
+ train.add_argument("--end-date")
61
+ train.add_argument(
62
+ "--skip-scoring",
63
+ action="store_true",
64
+ help="Skip scoring episodes that have no rewards yet.",
65
+ )
66
+
67
+ score = sub.add_parser("score", help="Score episodes but skip the policy update.")
68
+ score.add_argument("--agent-id", required=True)
69
+ score.add_argument("--task-id")
70
+ score.add_argument("--limit", type=_episode_limit, default=100)
71
+
72
+ show = sub.add_parser("task-policy", help="Print the active policy for an agent task.")
73
+ show.add_argument("--agent-id", required=True)
74
+ show.add_argument("--task-id", required=True)
75
+
76
+ init = sub.add_parser(
77
+ "task-policy-init",
78
+ help="Create and activate the initial policy for an agent task.",
79
+ )
80
+ init.add_argument("--agent-id", required=True)
81
+ init.add_argument("--task-id", required=True)
82
+ init.add_argument(
83
+ "--actions",
84
+ required=True,
85
+ help="Path to a JSON file containing a list of {id, description, parameters} objects.",
86
+ )
87
+
88
+ register = sub.add_parser(
89
+ "task-episode-register",
90
+ help="Register an agent task episode from a JSON file.",
91
+ )
92
+ register.add_argument("--agent-id", required=True)
93
+ register.add_argument("--task-id", required=True)
94
+ register.add_argument(
95
+ "--episode",
96
+ required=True,
97
+ help="Path to a JSON file containing an Episode object.",
98
+ )
99
+
100
+ return parser
101
+
102
+
103
+ def _cmd_agents_list(args: argparse.Namespace) -> int:
104
+ del args
105
+ agents = get_default_store().list_agents()
106
+ print(json.dumps([{"id": agent.id, "name": agent.name} for agent in agents], indent=2))
107
+ return 0
108
+
109
+
110
+ def _cmd_agent_tasks_list(args: argparse.Namespace) -> int:
111
+ tasks = get_default_store().list_agent_tasks(args.agent_id)
112
+ print(json.dumps([{"id": task.id, "name": task.name} for task in tasks], indent=2))
113
+ return 0
114
+
115
+
116
+ def _cmd_agents_episodes_count(args: argparse.Namespace) -> int:
117
+ count = get_default_store().count_episodes(
118
+ args.agent_id,
119
+ task_id=args.task_id,
120
+ full_only=True,
121
+ )
122
+ print(count)
123
+ return 0
124
+
125
+
126
+ def _cmd_agents_episodes_list(args: argparse.Namespace) -> int:
127
+ store = get_default_store()
128
+ episodes = store.query_episodes(
129
+ args.agent_id,
130
+ task_id=args.task_id,
131
+ limit=_MAX_EPISODES,
132
+ )
133
+ if not args.include_incomplete:
134
+ episodes = [episode for episode in episodes if episode.is_full]
135
+
136
+ payload = []
137
+ for episode in episodes[: args.limit]:
138
+ metrics = store.get_metric_results(episode.id, args.agent_id)
139
+ rewards = store.get_rewards_for_episode(episode.id, args.agent_id)
140
+ aggregate_rewards = [
141
+ reward for reward in rewards if reward.source == RewardSource.AGGREGATE
142
+ ]
143
+ aggregate_rewards.sort(key=lambda reward: reward.created_at, reverse=True)
144
+ completion = next(
145
+ (metric for metric in metrics if metric.metric == MetricName.TASK_COMPLETION),
146
+ None,
147
+ )
148
+ payload.append(
149
+ {
150
+ "episode": episode.to_dict(),
151
+ "score_breakdown": [metric.to_dict() for metric in metrics],
152
+ "final_reward": aggregate_rewards[0].value if aggregate_rewards else None,
153
+ "task_completion_quality": completion.to_dict() if completion else None,
154
+ }
155
+ )
156
+ print(json.dumps(payload, indent=2))
157
+ return 0
158
+
159
+
160
+ def _cmd_train(args: argparse.Namespace) -> int:
161
+ store = get_default_store()
162
+ task_ids = (
163
+ [args.task_id]
164
+ if args.task_id
165
+ else [task.id for task in store.list_agent_tasks(args.agent_id)]
166
+ )
167
+ selected_episodes = store.query_episodes(
168
+ args.agent_id,
169
+ task_id=args.task_id,
170
+ limit=args.limit,
171
+ start_date=args.start_date,
172
+ end_date=args.end_date,
173
+ )
174
+ episode_limits = Counter(episode.task_id for episode in selected_episodes)
175
+ runs: list[dict[str, Any]] = []
176
+ skipped: list[dict[str, str]] = []
177
+ for task_id in task_ids:
178
+ snapshot = store.get_active_policy(args.agent_id, task_id)
179
+ if snapshot is None:
180
+ skipped.append({"task_id": task_id, "reason": "no active policy"})
181
+ continue
182
+ episode_limit = episode_limits.get(task_id, 0)
183
+ if episode_limit == 0:
184
+ skipped.append({"task_id": task_id, "reason": "no episodes in selected batch"})
185
+ continue
186
+ policy = SoftmaxPolicy.from_snapshot(snapshot)
187
+ runner = LearningRunner(store=store, policy=policy)
188
+ run = runner.run_offline_batch(
189
+ args.agent_id,
190
+ task_id=task_id,
191
+ episode_limit=episode_limit,
192
+ start_date=args.start_date,
193
+ end_date=args.end_date,
194
+ score_missing=not args.skip_scoring,
195
+ )
196
+ runs.append(run.to_dict())
197
+
198
+ print(json.dumps({"agent_id": args.agent_id, "runs": runs, "skipped": skipped}, indent=2))
199
+ if runs:
200
+ return 0
201
+ print(
202
+ f"No task policies were trained for agent_id={args.agent_id!r}. "
203
+ "Check the skipped reasons and initialize missing task policies.",
204
+ file=sys.stderr,
205
+ )
206
+ return 2
207
+
208
+
209
+ def _cmd_score(args: argparse.Namespace) -> int:
210
+ store = get_default_store()
211
+ runner = LearningRunner(store=store)
212
+ episodes = store.query_episodes(
213
+ args.agent_id,
214
+ task_id=args.task_id,
215
+ limit=args.limit,
216
+ )
217
+ scored = 0
218
+ for episode in episodes:
219
+ existing = store.get_rewards_for_episode(episode.id, args.agent_id)
220
+ if existing:
221
+ continue
222
+ runner.score_and_record(episode)
223
+ scored += 1
224
+ print(json.dumps({"episodes_seen": len(episodes), "newly_scored": scored}, indent=2))
225
+ return 0
226
+
227
+
228
+ def _policy_payload(snapshot: PolicySnapshot) -> dict[str, Any]:
229
+ payload = snapshot.to_dict()
230
+ policy = SoftmaxPolicy.from_snapshot(snapshot)
231
+ payload["action_probabilities"] = {
232
+ action.id: probability
233
+ for action, probability in zip(snapshot.actions, policy.probabilities())
234
+ }
235
+ return payload
236
+
237
+
238
+ def _policy_difference(
239
+ current: PolicySnapshot, previous: PolicySnapshot | None
240
+ ) -> dict[str, Any] | None:
241
+ if previous is None:
242
+ return None
243
+ current_payload = _policy_payload(current)
244
+ previous_payload = _policy_payload(previous)
245
+ action_ids = sorted(
246
+ set(current_payload["action_probabilities"])
247
+ | set(previous_payload["action_probabilities"])
248
+ )
249
+ return {
250
+ "version": current.version - previous.version,
251
+ "baseline": current.baseline - previous.baseline,
252
+ "episodes_seen": current.episodes_seen - previous.episodes_seen,
253
+ "updates_applied": current.updates_applied - previous.updates_applied,
254
+ "logits": {
255
+ action_id: current.logits.get(action_id, 0.0)
256
+ - previous.logits.get(action_id, 0.0)
257
+ for action_id in action_ids
258
+ },
259
+ "action_probabilities": {
260
+ action_id: current_payload["action_probabilities"].get(action_id, 0.0)
261
+ - previous_payload["action_probabilities"].get(action_id, 0.0)
262
+ for action_id in action_ids
263
+ },
264
+ }
265
+
266
+
267
+ def _cmd_show_task_policy(args: argparse.Namespace) -> int:
268
+ store = get_default_store()
269
+ snapshot = store.get_active_policy(args.agent_id, args.task_id)
270
+ if snapshot is None:
271
+ print(
272
+ f"No active policy found for agent_id={args.agent_id!r}, "
273
+ f"task_id={args.task_id!r}.",
274
+ file=sys.stderr,
275
+ )
276
+ return 2
277
+ previous = next(
278
+ (
279
+ policy
280
+ for policy in store.list_policies(args.agent_id, args.task_id)
281
+ if policy.id != snapshot.id
282
+ ),
283
+ None,
284
+ )
285
+ print(
286
+ json.dumps(
287
+ {
288
+ "current_policy": _policy_payload(snapshot),
289
+ "previous_policy": _policy_payload(previous) if previous else None,
290
+ "difference": _policy_difference(snapshot, previous),
291
+ },
292
+ indent=2,
293
+ )
294
+ )
295
+ return 0
296
+
297
+
298
+ def _cmd_init_task_policy(args: argparse.Namespace) -> int:
299
+ store = get_default_store()
300
+ if store.get_active_policy(args.agent_id, args.task_id) is not None:
301
+ print(
302
+ f"An active policy already exists for agent_id={args.agent_id!r}, "
303
+ f"task_id={args.task_id!r}.",
304
+ file=sys.stderr,
305
+ )
306
+ return 2
307
+ try:
308
+ with open(args.actions, "r", encoding="utf-8") as actions_file:
309
+ action_payloads = json.load(actions_file)
310
+ except (OSError, json.JSONDecodeError) as exc:
311
+ print(f"Unable to read --actions file: {exc}", file=sys.stderr)
312
+ return 2
313
+ if not isinstance(action_payloads, list) or not action_payloads:
314
+ print("--actions file must contain a non-empty JSON list", file=sys.stderr)
315
+ return 2
316
+ try:
317
+ actions = [Action.from_dict(item) for item in action_payloads]
318
+ except (KeyError, TypeError, ValueError) as exc:
319
+ print(f"Invalid action definition: {exc}", file=sys.stderr)
320
+ return 2
321
+ policy = SoftmaxPolicy.from_actions(
322
+ actions,
323
+ agent_id=args.agent_id,
324
+ task_id=args.task_id,
325
+ )
326
+ snapshot = policy.snapshot()
327
+ store.store_policy(snapshot)
328
+ print(json.dumps(_policy_payload(snapshot), indent=2))
329
+ return 0
330
+
331
+
332
+ def _cmd_register_task_episode(args: argparse.Namespace) -> int:
333
+ try:
334
+ with open(args.episode, "r", encoding="utf-8") as episode_file:
335
+ payload = json.load(episode_file)
336
+ except (OSError, json.JSONDecodeError) as exc:
337
+ print(f"Unable to read --episode file: {exc}", file=sys.stderr)
338
+ return 2
339
+ if not isinstance(payload, dict):
340
+ print("--episode file must contain a JSON object", file=sys.stderr)
341
+ return 2
342
+
343
+ payload = dict(payload)
344
+ for field, expected in (("agent_id", args.agent_id), ("task_id", args.task_id)):
345
+ actual = payload.get(field)
346
+ if actual is not None and actual != expected:
347
+ print(
348
+ f"Episode {field}={actual!r} does not match --{field.replace('_', '-')}={expected!r}.",
349
+ file=sys.stderr,
350
+ )
351
+ return 2
352
+ payload[field] = expected
353
+ payload.setdefault("id", str(uuid.uuid4()))
354
+
355
+ try:
356
+ episode = Episode.from_dict(payload)
357
+ except (KeyError, TypeError, ValueError) as exc:
358
+ print(f"Invalid episode definition: {exc}", file=sys.stderr)
359
+ return 2
360
+
361
+ get_default_store().store_episode(episode)
362
+ print(json.dumps(episode.to_dict(), indent=2))
363
+ return 0
364
+
365
+
366
+ def main(argv: list[str] | None = None) -> int:
367
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s - %(message)s")
368
+ parser = _build_arg_parser()
369
+ args = parser.parse_args(argv)
370
+ dispatch = {
371
+ "list": _cmd_agents_list,
372
+ "tasks-list": _cmd_agent_tasks_list,
373
+ "task-episodes-count": _cmd_agents_episodes_count,
374
+ "task-episodes-list": _cmd_agents_episodes_list,
375
+ "train": _cmd_train,
376
+ "score": _cmd_score,
377
+ "task-policy": _cmd_show_task_policy,
378
+ "task-policy-init": _cmd_init_task_policy,
379
+ "task-episode-register": _cmd_register_task_episode,
380
+ }
381
+ handler = dispatch[args.command]
382
+ return handler(args)
383
+
384
+
385
+ if __name__ == "__main__": # pragma: no cover
386
+ sys.exit(main())