summaryrefslogtreecommitdiff
path: root/worldalign/energy_functional_infer.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/energy_functional_infer.py')
-rw-r--r--worldalign/energy_functional_infer.py262
1 files changed, 262 insertions, 0 deletions
diff --git a/worldalign/energy_functional_infer.py b/worldalign/energy_functional_infer.py
new file mode 100644
index 0000000..6864011
--- /dev/null
+++ b/worldalign/energy_functional_infer.py
@@ -0,0 +1,262 @@
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import torch
+import torch.nn.functional as F
+from transformers import AutoModelForCausalLM, AutoTokenizer
+
+from .common import dtype_for_device, read_json, seed_everything, write_json
+from .energy import (
+ projection_quantile_target,
+ relation_field_energy,
+ retrieval_metrics,
+ sliced_distribution_energy,
+ standardized_relation,
+)
+from .extract_text import mean_pool
+from .io import load_feature_pair, select_rows
+from .models import load_prefix
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--manifest", default="artifacts/manifest.json")
+ parser.add_argument("--vision", default="artifacts/vision.pt")
+ parser.add_argument("--text", default="artifacts/text.pt")
+ parser.add_argument(
+ "--text-orbits", default="artifacts/text_orbits_qwen0p5b.pt"
+ )
+ parser.add_argument("--prefix", default="artifacts/prefix.pt")
+ parser.add_argument("--split", choices=["val", "test"], default="val")
+ parser.add_argument("--samples", type=int, default=128)
+ parser.add_argument("--steps", type=int, default=60)
+ parser.add_argument("--refresh-steps", type=int, default=20)
+ parser.add_argument("--max-new-tokens", type=int, default=24)
+ parser.add_argument("--lr", type=float, default=0.03)
+ parser.add_argument("--projections", type=int, default=128)
+ parser.add_argument("--relation-weight", type=float, default=1.0)
+ parser.add_argument("--conditional-weight", type=float, default=0.08)
+ parser.add_argument("--distribution-weight", type=float, default=80.0)
+ parser.add_argument("--functional-weight", type=float, default=10.0)
+ parser.add_argument("--device", default="cuda:1")
+ parser.add_argument("--seed", type=int, default=20260729)
+ parser.add_argument(
+ "--output", default="artifacts/energy_functional_val.pt"
+ )
+ parser.add_argument(
+ "--metrics-output", default="artifacts/energy_functional_val.json"
+ )
+ return parser.parse_args()
+
+
+@torch.no_grad()
+def refresh_functional_anchor(
+ particles: torch.Tensor,
+ prefix: torch.nn.Module,
+ lm: AutoModelForCausalLM,
+ tokenizer: AutoTokenizer,
+ dtype: torch.dtype,
+ max_new_tokens: int,
+) -> tuple[torch.Tensor, list[str]]:
+ captions: list[str] = []
+ for chunk in particles.split(16):
+ prefix_embedding = prefix(chunk).to(dtype)
+ attention = torch.ones(
+ prefix_embedding.shape[:2],
+ dtype=torch.long,
+ device=particles.device,
+ )
+ output = lm.generate(
+ inputs_embeds=prefix_embedding,
+ attention_mask=attention,
+ max_new_tokens=max_new_tokens,
+ do_sample=False,
+ eos_token_id=tokenizer.eos_token_id,
+ pad_token_id=tokenizer.pad_token_id,
+ )
+ captions.extend(
+ tokenizer.batch_decode(output, skip_special_tokens=True)
+ )
+ anchor_parts: list[torch.Tensor] = []
+ for start in range(0, len(captions), 32):
+ tokens = tokenizer(
+ captions[start : start + 32],
+ padding=True,
+ truncation=True,
+ max_length=64,
+ return_tensors="pt",
+ )
+ tokens = {
+ key: value.to(particles.device)
+ for key, value in tokens.items()
+ }
+ result = lm(
+ **tokens, output_hidden_states=True, return_dict=True
+ )
+ anchor_parts.append(
+ mean_pool(
+ result.hidden_states[-1], tokens["attention_mask"]
+ ).float()
+ )
+ return F.normalize(torch.cat(anchor_parts), dim=-1), captions
+
+
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ manifest = read_json(args.manifest)
+ vision, text, vision_lookup, _ = load_feature_pair(
+ args.vision, args.text
+ )
+ orbit_state = torch.load(
+ args.text_orbits, map_location="cpu", weights_only=False
+ )
+ orbit_lookup = {
+ int(row): index for index, row in enumerate(orbit_state["rows"])
+ }
+ orbit_mean = F.normalize(
+ orbit_state["features"].float().mean(1), dim=-1
+ )
+ rows = manifest[args.split][: args.samples]
+ visual = select_rows(
+ vision["features"], vision_lookup, rows
+ ).to(args.device)
+ paired_text = select_rows(
+ orbit_mean, orbit_lookup, rows
+ ).to(args.device)
+ text_population = select_rows(
+ orbit_mean, orbit_lookup, manifest["text_only_train"]
+ ).to(args.device)
+
+ prefix, prefix_state = load_prefix(args.prefix, args.device)
+ if prefix.semantic_dim != text_population.shape[-1]:
+ raise ValueError("Prefix and text-orbit dimensions do not match")
+ tokenizer = AutoTokenizer.from_pretrained(text["model"])
+ if tokenizer.pad_token_id is None:
+ tokenizer.pad_token = tokenizer.eos_token
+ dtype = dtype_for_device(args.device)
+ lm = AutoModelForCausalLM.from_pretrained(
+ text["model"], torch_dtype=dtype
+ ).to(args.device)
+ lm.eval()
+ for parameter in lm.parameters():
+ parameter.requires_grad_(False)
+ for parameter in prefix.parameters():
+ parameter.requires_grad_(False)
+
+ generator = torch.Generator(device=args.device).manual_seed(args.seed)
+ initial_indices = torch.randperm(
+ len(text_population),
+ generator=generator,
+ device=args.device,
+ )[: len(visual)]
+ particles = torch.nn.Parameter(text_population[initial_indices].clone())
+ initial_particles = particles.detach().cpu()
+ directions, target_quantiles = projection_quantile_target(
+ text_population,
+ len(particles),
+ args.projections,
+ generator,
+ )
+ visual_relation, visual_standardized = standardized_relation(visual)
+ optimizer = torch.optim.Adam([particles], lr=args.lr)
+ functional_anchor, functional_captions = refresh_functional_anchor(
+ particles.detach(),
+ prefix,
+ lm,
+ tokenizer,
+ dtype,
+ args.max_new_tokens,
+ )
+
+ history: list[dict] = []
+ refresh_captions: dict[int, list[str]] = {
+ 0: functional_captions[:25]
+ }
+ for step in range(args.steps + 1):
+ if step > 0 and step % args.refresh_steps == 0:
+ functional_anchor, functional_captions = (
+ refresh_functional_anchor(
+ particles.detach(),
+ prefix,
+ lm,
+ tokenizer,
+ dtype,
+ args.max_new_tokens,
+ )
+ )
+ refresh_captions[step] = functional_captions[:25]
+ relation, conditional = relation_field_energy(
+ visual_relation, visual_standardized, particles
+ )
+ distribution = sliced_distribution_energy(
+ particles, directions, target_quantiles
+ )
+ functional = (
+ 1
+ - (
+ F.normalize(particles, dim=-1)
+ * functional_anchor.detach()
+ ).sum(-1)
+ ).mean()
+ loss = (
+ args.relation_weight * relation
+ + args.conditional_weight * conditional
+ + args.distribution_weight * distribution
+ + args.functional_weight * functional
+ )
+ if step % 10 == 0 or step == args.steps:
+ record = {
+ "step": step,
+ "total": float(loss.detach()),
+ "relation": float(relation.detach()),
+ "conditional": float(conditional.detach()),
+ "distribution": float(distribution.detach()),
+ "functional": float(functional.detach()),
+ "paired_evaluation_only": retrieval_metrics(
+ particles.detach(), paired_text
+ ),
+ }
+ history.append(record)
+ print(json.dumps(record))
+ if step == args.steps:
+ break
+ optimizer.zero_grad(set_to_none=True)
+ loss.backward()
+ torch.nn.utils.clip_grad_norm_([particles], 2.0)
+ optimizer.step()
+ with torch.no_grad():
+ particles.copy_(F.normalize(particles, dim=-1))
+
+ result = {
+ "protocol": (
+ "No image-text pair and no cross-modal parameter is used. The "
+ "functional language energy is a frozen decode-reencode cycle "
+ "through a text-only prefix interface and frozen Qwen."
+ ),
+ "mode": "functional_cycle_language_latent_particles",
+ "split": args.split,
+ "rows": rows,
+ "args": vars(args),
+ "history": history,
+ "prefix_training": prefix_state["training"],
+ "refresh_caption_examples": refresh_captions,
+ }
+ state = {
+ **result,
+ "initial_particles": initial_particles,
+ "final_particles": F.normalize(
+ particles.detach(), dim=-1
+ ).cpu(),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ write_json(args.metrics_output, result)
+ print(f"Wrote {args.output} and {args.metrics_output}")
+
+
+if __name__ == "__main__":
+ main()