my-cool-model / lib.py
reshma0639's picture
Upload folder using huggingface_hub (part 15)
b0ba444 verified
Raw
History Blame Contribute Delete
15.9 kB
"""
mosketch_lib.py — all non-Qwen, non-CLI logic for the sketch deformation
pipeline in one place: loading SVG/semantic/trajectory/caption files,
building meshes, running ARAP, and rendering results.
This file has NO command-line interface and makes NO calls to Qwen or any
other model — it's a pure library, imported by mosketch_pipeline.py.
Consolidated from (all logic ported, not rewritten): arap_2d.py,
trajectory.py, captions.py, and the mesh/handle-selection functions from
auto_pipeline.py (auto_select_handles_deduped and its dependencies —
auto_pipeline.py's own moving/static classifier and rigidity heuristic
were NOT ported here, since both were confirmed unused by the real
pipeline and the rigidity heuristic was confirmed broken on real data).
"""
import re
import numpy as np
import scipy.sparse as sp
import scipy.sparse.linalg as spla
from scipy.spatial import Delaunay, cKDTree
from sklearn.cluster import KMeans
from svgpathtools import svg2paths
# =============================================================================
# SVG / semantic loading
# =============================================================================
def load_strokes_from_svg(svg_path, samples_per_stroke=12):
paths, _ = svg2paths(svg_path)
strokes = []
for p in paths:
if len(p) == 0:
continue
ts = np.linspace(0, 1, samples_per_stroke)
pts = np.array([[p.point(t).real, p.point(t).imag] for t in ts])
strokes.append(pts)
if not strokes:
raise ValueError(f"No usable paths found in {svg_path}")
return strokes
def load_semantic_assignments(path):
"""Returns dict: object_name -> list of int stroke indices."""
assignments = {}
with open(path) as f:
for line in f:
line = line.rstrip("\n")
if not line.strip():
continue
name, idx_str = line.split("\t")
indices = [int(i) for i in idx_str.split(",") if i.strip() != ""]
assignments[name] = indices
return assignments
def filter_strokes(strokes, indices):
return [strokes[i] for i in indices]
def flatten_strokes(strokes):
"""Returns (all_points (N,2), stroke_slices [(start,end), ...])."""
all_points = []
stroke_slices = []
offset = 0
for stroke in strokes:
n = len(stroke)
stroke_slices.append((offset, offset + n))
all_points.append(stroke)
offset += n
return np.vstack(all_points), stroke_slices
def load_object(name, svg_path, semantic_path):
"""Convenience: load one named object's flattened points + slices directly."""
strokes = load_strokes_from_svg(svg_path)
assignments = load_semantic_assignments(semantic_path)
strokes = filter_strokes(strokes, assignments[name])
return flatten_strokes(strokes)
# =============================================================================
# Mesh construction (dedup + Delaunay)
# =============================================================================
def deduplicate_points(points, tol=1e-3):
"""
Merges points within `tol` of each other into a single mesh vertex.
Necessary because strokes meeting at a shared joint often have
exactly (or near-exactly) coincident coordinates — Delaunay/QHull can
leave one copy of a duplicate pair disconnected from every triangle,
which makes the ARAP Laplacian singular.
Returns:
unique_points: (M, 2) deduplicated point positions.
point_to_unique: (N,) int array mapping each original point index
to its merged vertex index in unique_points.
"""
N = len(points)
tree = cKDTree(points)
parent = list(range(N))
def find(x):
while parent[x] != x:
parent[x] = parent[parent[x]]
x = parent[x]
return x
def union(a, b):
ra, rb = find(a), find(b)
if ra != rb:
parent[ra] = rb
pairs = tree.query_pairs(r=tol)
for a, b in pairs:
union(a, b)
roots = sorted(set(find(i) for i in range(N)))
root_to_unique = {r: i for i, r in enumerate(roots)}
point_to_unique = np.array([root_to_unique[find(i)] for i in range(N)])
unique_points = np.zeros((len(roots), 2))
counts = np.zeros(len(roots))
for i in range(N):
u = point_to_unique[i]
unique_points[u] += points[i]
counts[u] += 1
unique_points /= counts[:, None]
return unique_points, point_to_unique
def build_mesh(points):
"""Delaunay-triangulates the (already deduplicated) point cloud."""
tri = Delaunay(points)
edges = set()
for simplex in tri.simplices:
for i in range(3):
a, b = simplex[i], simplex[(i + 1) % 3]
edges.add((min(a, b), max(a, b)))
return tri, np.array(sorted(edges))
def nearest_mesh_vertex(unique_points, target):
dists = np.linalg.norm(unique_points - np.array(target), axis=1)
return int(dists.argmin())
# =============================================================================
# Automatic joint/handle selection (KMeans + collision-safe dedup)
# =============================================================================
def auto_select_handles(strokes, k=4, seed=0):
"""
Runs KMeans on stroke centroids to get k joints. The joint nearest the
object's own centroid becomes the anchor; all others become handles.
Returns RAW (non-deduplicated) joint positions — use
auto_select_handles_deduped for the mesh-safe version.
"""
centroids = np.array([s.mean(axis=0) for s in strokes])
km = KMeans(n_clusters=k, n_init=10, random_state=seed).fit(centroids)
joints = km.cluster_centers_
object_centroid = centroids.mean(axis=0)
dists = np.linalg.norm(joints - object_centroid, axis=1)
anchor_idx = int(dists.argmin())
handle_idxs = [i for i in range(k) if i != anchor_idx]
return joints, anchor_idx, handle_idxs
def auto_select_handles_deduped(strokes, unique_points, k=4, seed=0):
"""
Like auto_select_handles, but maps each joint to its nearest mesh vertex
first and DROPS any joint that collides with an already-claimed vertex.
Necessary for small objects (e.g. a frisbee) where KMeans can find k
"distinct" cluster centers that all round to the same underlying mesh
point after deduplication — without this, two different joint names
would silently control the same physical vertex, and whichever target
arrives last would overwrite the other with no error.
Returns:
joints: (M, 2) deduped joint positions, M <= k.
anchor_idx: index into `joints`.
handle_idxs: list of indices into `joints`.
joint_mesh_indices: list of mesh vertex indices, one per joint,
guaranteed unique.
"""
joints_raw, _, _ = auto_select_handles(strokes, k=k, seed=seed)
seen_vertices = {}
kept_joints = []
kept_mesh_indices = []
for j in joints_raw:
mv = nearest_mesh_vertex(unique_points, j)
if mv in seen_vertices:
continue
seen_vertices[mv] = True
kept_joints.append(j)
kept_mesh_indices.append(mv)
kept_joints = np.array(kept_joints)
object_centroid = kept_joints.mean(axis=0)
dists = np.linalg.norm(kept_joints - object_centroid, axis=1)
anchor_idx = int(dists.argmin())
handle_idxs = [i for i in range(len(kept_joints)) if i != anchor_idx]
dropped = k - len(kept_joints)
if dropped > 0:
print(f" NOTE: {dropped} of {k} joints collided onto an already-used mesh vertex "
f"and were dropped — object only supports {len(kept_joints)} independent handles")
return kept_joints, anchor_idx, handle_idxs, kept_mesh_indices
def object_bbox_size(points):
return np.linalg.norm(points.max(axis=0) - points.min(axis=0))
def build_stroke_geometry_text(strokes, n_points=2):
"""
Formats an object's actual stroke points as compact text — e.g.
"stroke_0: (12,45) -> (22,53)" per stroke. Built specifically to give
Qwen the object's real geometry in TEXT form, not just a rendered
image: there's real literature suggesting multimodal LLMs can
under-attend to image tokens relative to text, so this is a
complementary, text-native representation of the same information the
image already shows, not a replacement for it.
n_points=2 (start+end of each stroke only) was chosen after measuring
real cost on dog3's 64-stroke dog object: 2 points/stroke = ~570
tokens, 4 points/stroke = ~990 tokens, full 12 points/stroke = ~2650
tokens — given this pipeline has already hit real OOM/latency limits
this session, the sparsest still-useful setting was chosen
deliberately, not the richest one.
"""
lines = []
for i, stroke in enumerate(strokes):
idxs = np.linspace(0, len(stroke) - 1, n_points).astype(int)
pts = stroke[idxs]
pts_str = " -> ".join(f"({x:.0f},{y:.0f})" for x, y in pts)
lines.append(f" stroke_{i}: {pts_str}")
return "\n".join(lines)
# =============================================================================
# ARAP solver (Sorkine & Alexa 2007, local-global, uniform edge weights)
# =============================================================================
def arap_deform(points, edges, handle_indices, handle_targets, iterations=10):
"""
points: (N, 2) rest-pose positions.
edges: (E, 2) int array of vertex index pairs.
handle_indices: list of point indices with fixed target positions.
handle_targets: (len(handle_indices), 2) target positions.
Returns: deformed points, (N, 2).
"""
N = len(points)
neighbors = {i: [] for i in range(N)}
for a, b in edges:
neighbors[a].append(b)
neighbors[b].append(a)
W = sp.lil_matrix((N, N))
for a, b in edges:
W[a, b] = 1.0
W[b, a] = 1.0
W = W.tocsr()
deg = np.array(W.sum(axis=1)).flatten()
L = sp.diags(deg) - W
p_prime = points.copy()
for idx, target in zip(handle_indices, handle_targets):
p_prime[idx] = target
L_constrained = L.tolil()
for idx in handle_indices:
L_constrained.rows[idx] = [idx]
L_constrained.data[idx] = [1.0]
L_constrained = L_constrained.tocsr()
solver = spla.factorized(L_constrained.tocsc())
for it in range(iterations):
rotations = np.zeros((N, 2, 2))
for i in range(N):
js = neighbors[i]
if not js:
rotations[i] = np.eye(2)
continue
P = np.array([points[i] - points[j] for j in js]).T
Pp = np.array([p_prime[i] - p_prime[j] for j in js]).T
S = P @ Pp.T
U, _, Vt = np.linalg.svd(S)
R = Vt.T @ U.T
if np.linalg.det(R) < 0:
Vt[-1, :] *= -1
R = Vt.T @ U.T
rotations[i] = R
bx = np.zeros(N)
by = np.zeros(N)
for a, b in edges:
e_rest = points[a] - points[b]
contrib = 0.5 * (rotations[a] + rotations[b]) @ e_rest
bx[a] += contrib[0]; by[a] += contrib[1]
bx[b] -= contrib[0]; by[b] -= contrib[1]
for idx, target in zip(handle_indices, handle_targets):
bx[idx] = target[0]
by[idx] = target[1]
x_new = solver(bx)
y_new = solver(by)
p_prime = np.stack([x_new, y_new], axis=1)
return p_prime
# =============================================================================
# Trajectory parsing (real _traj.txt files)
# =============================================================================
DEFAULT_KEYFRAME_INDICES = (0, 4, 9, 14, 19)
def load_trajectories(traj_path, keyframe_indices=DEFAULT_KEYFRAME_INDICES):
"""
File format: tab-separated, one object per line —
object_name\\t[x,y,w,h],[x,y,w,h],... (20 frames)
Returns dict: object_name -> list of (x,y,w,h), sampled at keyframe_indices.
"""
trajectories = {}
with open(traj_path) as f:
for line_no, line in enumerate(f, 1):
line = line.rstrip("\n")
if not line.strip():
continue
if "\t" not in line:
raise ValueError(
f"{traj_path} line {line_no}: expected tab-separated "
f"'object_name\\t[x,y,w,h],...' format, got: {line[:80]!r}"
)
name, boxes_str = line.split("\t", 1)
box_strs = re.findall(r"\[([^\]]+)\]", boxes_str)
if not box_strs:
raise ValueError(f"{traj_path} line {line_no}: no [x,y,w,h] boxes found for '{name}'")
boxes = [tuple(float(v) for v in b.split(",")) for b in box_strs]
max_idx = max(keyframe_indices)
if len(boxes) <= max_idx:
raise ValueError(
f"{traj_path} line {line_no}: object '{name}' has only {len(boxes)} "
f"frames, need at least {max_idx + 1} for keyframe indices {keyframe_indices}"
)
trajectories[name] = [boxes[i] for i in keyframe_indices]
return trajectories
def bbox_deltas(sampled_boxes):
"""Returns (dx, dy) numpy arrays, offset from the first sampled keyframe."""
xs = np.array([b[0] for b in sampled_boxes], dtype=float)
ys = np.array([b[1] for b in sampled_boxes], dtype=float)
return xs - xs[0], ys - ys[0]
def movement_magnitude(sampled_boxes):
dx, dy = bbox_deltas(sampled_boxes)
return float(np.hypot(dx[-1], dy[-1]))
# =============================================================================
# Caption parsing (real caption.txt files)
# =============================================================================
def load_captions(caption_path):
"""
File format (CONFIRMED against the real 60-sketch caption.txt):
tab-separated, one sketch per line — sketch_name\\tcaption text.
Returns dict: sketch_name -> caption string.
"""
captions = {}
with open(caption_path) as f:
for line_no, line in enumerate(f, 1):
line = line.rstrip("\n")
if not line.strip():
continue
if "\t" not in line:
raise ValueError(
f"{caption_path} line {line_no}: expected tab-separated "
f"'sketch_name\\tcaption' format, got: {line[:80]!r}"
)
name, caption = line.split("\t", 1)
captions[name.strip()] = caption.strip()
return captions
def get_caption(caption_path, sketch_name):
captions = load_captions(caption_path)
if sketch_name not in captions:
raise KeyError(
f"'{sketch_name}' not found in {caption_path}. "
f"Available: {list(captions.keys())[:10]}{'...' if len(captions) > 10 else ''}"
)
return captions[sketch_name]
# =============================================================================
# Rendering
# =============================================================================
def render(points, stroke_slices, mesh_edges=None, handles=None, out_path="out.png", title=""):
import matplotlib.pyplot as plt
fig, ax = plt.subplots(figsize=(6, 6))
if mesh_edges is not None:
for a, b in mesh_edges:
ax.plot([points[a, 0], points[b, 0]], [points[a, 1], points[b, 1]],
color="#D9DEE2", linewidth=0.5, zorder=1)
for start, end in stroke_slices:
seg = points[start:end]
ax.plot(seg[:, 0], seg[:, 1], color="black", linewidth=1.3, zorder=2)
if handles is not None:
ax.scatter(points[handles, 0], points[handles, 1], color="#3D7A5C", s=80, zorder=5)
ax.invert_yaxis()
ax.set_aspect("equal")
ax.set_title(title)
fig.savefig(out_path, dpi=150, bbox_inches="tight")
plt.close(fig)