make adaptivelr resolution aware to balance adjusments
This commit is contained in:
parent
1b71546a62
commit
459e794134
|
|
@ -920,11 +920,17 @@ class SDTrainer(BaseSDTrainProcess):
|
|||
# (flow-matching schedulers included) — normalize to [0, 1] for the bucket index.
|
||||
num_train_timesteps = float(getattr(self.train_config, 'num_train_timesteps', 1000) or 1000)
|
||||
for idx, file_item in enumerate(batch.file_items):
|
||||
# nominal configured resolution (e.g. 256/512/1024), NOT the post-aspect-crop
|
||||
# dimensions — with bucketing, crop_width/crop_height vary per aspect ratio
|
||||
# within the same resolution tier (256/288/336/... are all "256"), which would
|
||||
# fragment one tier into several and defeat the point of grouping by it.
|
||||
res = int(getattr(getattr(file_item, 'dataset_config', None), 'resolution', 0) or 0)
|
||||
self.loss_watch.observe(
|
||||
epoch=window_idx,
|
||||
item_key=file_item.path,
|
||||
timestep=float(ts_detached[idx].item()) / num_train_timesteps,
|
||||
loss=float(loss_detached[idx].item()),
|
||||
resolution=res,
|
||||
)
|
||||
except Exception as e:
|
||||
print_acc(f"[adaptive-lr] observe failed: {e}")
|
||||
|
|
|
|||
|
|
@ -96,9 +96,15 @@ class PerImageAdaptiveLR:
|
|||
# instead (BaseSDTrainProcess passes print_acc) — plain print() otherwise.
|
||||
self._log = log_fn or print
|
||||
|
||||
self._records: list[tuple[str, int, int, float]] = [] # (key, epoch, bucket, loss)
|
||||
self._bsum = [0.0] * _N_BUCKETS
|
||||
self._bcnt = [0] * _N_BUCKETS
|
||||
# bucket key is now (timestep_bucket, resolution) — resolution is the trained long-edge
|
||||
# pixel size (e.g. 512/768/1024). Pooling residuals against a bucket mean that ignores
|
||||
# resolution would let a structurally harder resolution (more detail to predict, higher
|
||||
# loss for EVERYONE at that size) get misread as "this image is hard", which can throttle
|
||||
# a perfectly fine image's low-res copies just because its high-res copy is intrinsically
|
||||
# tougher — or mask a genuinely bad caption whose effect is only visible at one size.
|
||||
self._records: list[tuple[str, int, int, int, float]] = [] # (key, epoch, tbucket, res, loss)
|
||||
self._bsum: dict[tuple[int, int], float] = {}
|
||||
self._bcnt: dict[tuple[int, int], int] = {}
|
||||
|
||||
self._mult: dict[str, float] = {}
|
||||
self.verdicts: dict[str, str] = {}
|
||||
|
|
@ -125,16 +131,21 @@ class PerImageAdaptiveLR:
|
|||
"""Current loss multiplier for an item (looked up before the loss is scaled)."""
|
||||
return self._mult.get(str(item_key), 1.0)
|
||||
|
||||
def observe(self, *, epoch: int, item_key: str, timestep: float, loss: float) -> None:
|
||||
"""Record one image's loss for this step. Never raises into the training loop."""
|
||||
def observe(self, *, epoch: int, item_key: str, timestep: float, loss: float, resolution: int = 0) -> None:
|
||||
"""Record one image's loss for this step. `resolution` should be the trained long-edge
|
||||
pixel size (e.g. max(crop_width, crop_height)) — pass 0 (default) if unknown, which just
|
||||
pools everything into one resolution bucket, same as the old behavior. Never raises into
|
||||
the training loop."""
|
||||
try:
|
||||
key = str(item_key)
|
||||
t = float(timestep)
|
||||
loss = float(loss)
|
||||
res = int(resolution)
|
||||
b = self._bucket(t)
|
||||
self._bsum[b] += loss
|
||||
self._bcnt[b] += 1
|
||||
self._records.append((key, int(epoch), b, loss))
|
||||
bucket_key = (b, res)
|
||||
self._bsum[bucket_key] = self._bsum.get(bucket_key, 0.0) + loss
|
||||
self._bcnt[bucket_key] = self._bcnt.get(bucket_key, 0) + 1
|
||||
self._records.append((key, int(epoch), b, res, loss))
|
||||
except Exception as e:
|
||||
self._log(f"[adaptive-lr] observe failed ({e})")
|
||||
|
||||
|
|
@ -147,14 +158,16 @@ class PerImageAdaptiveLR:
|
|||
if epoch <= self.warmup_epochs or (not self._records and not self._restored_residuals):
|
||||
return self.verdicts
|
||||
|
||||
# Per-timestep-bucket means over the WHOLE run so far (order-independent), recomputed
|
||||
# fresh each boundary — matches the validated offline analyzer, not a live running mean.
|
||||
bmean = [self._bsum[i] / self._bcnt[i] if self._bcnt[i] else 0.0 for i in range(_N_BUCKETS)]
|
||||
# Per-(timestep-bucket, resolution) means over the WHOLE run so far (order-independent),
|
||||
# recomputed fresh each boundary — matches the validated offline analyzer, not a live
|
||||
# running mean. Keying by resolution too means a structurally harder/easier resolution
|
||||
# gets its own baseline instead of dragging every image's pooled residual with it.
|
||||
bmean = {bk: self._bsum[bk] / self._bcnt[bk] for bk in self._bsum if self._bcnt[bk]}
|
||||
|
||||
# per-key, per-epoch mean residual
|
||||
by_key_epoch: dict[str, dict[int, list[float]]] = {}
|
||||
for key, ep, b, loss in self._records:
|
||||
by_key_epoch.setdefault(key, {}).setdefault(ep, []).append(loss - bmean[b])
|
||||
for key, ep, b, res, loss in self._records:
|
||||
by_key_epoch.setdefault(key, {}).setdefault(ep, []).append(loss - bmean.get((b, res), 0.0))
|
||||
|
||||
per_key_epoch_mean: dict[str, dict[int, float]] = {
|
||||
key: {ep: sum(vals) / len(vals) for ep, vals in eps.items()}
|
||||
|
|
@ -259,16 +272,59 @@ class PerImageAdaptiveLR:
|
|||
for v in self.verdicts.values():
|
||||
counts[v] = counts.get(v, 0) + 1
|
||||
summary = ", ".join(f"{v}={c}" for v, c in sorted(counts.items()))
|
||||
self._log(f"[adaptive-lr] window {epoch}: {len(self.verdicts)} image(s) tracked — {summary}")
|
||||
|
||||
# Per-resolution raw average loss (marginalized over timestep bucket) — purely a
|
||||
# diagnostic, doesn't feed classification. Per-image verdicts can't tell you if a
|
||||
# WHOLE resolution tier is systematically worse (VRAM/precision quirk, a bucketing
|
||||
# bug, etc. rather than any individual image's data being bad) — this can. Appended
|
||||
# to the same line as the summary rather than a separate log line to keep it terse.
|
||||
res_sum: dict[int, float] = {}
|
||||
res_cnt: dict[int, int] = {}
|
||||
for (b, res), s in self._bsum.items():
|
||||
res_sum[res] = res_sum.get(res, 0.0) + s
|
||||
res_cnt[res] = res_cnt.get(res, 0) + self._bcnt[(b, res)]
|
||||
res_suffix = ""
|
||||
if len(res_sum) > 1:
|
||||
res_avgs = ", ".join(
|
||||
f"{res}px={res_sum[res] / res_cnt[res]:.4f}"
|
||||
for res in sorted(res_sum) if res_cnt[res]
|
||||
)
|
||||
res_suffix = f" — avg loss by res: {res_avgs}"
|
||||
|
||||
self._log(f"[adaptive-lr] window {epoch}: {len(self.verdicts)} image(s) tracked — "
|
||||
f"{summary}{res_suffix}")
|
||||
|
||||
added = self._confirmed_stuck - self._last_reported_stuck
|
||||
removed = self._last_reported_stuck - self._confirmed_stuck
|
||||
if added:
|
||||
names = ", ".join(os.path.basename(k) for k in sorted(added))
|
||||
self._log(f"[adaptive-lr] window {epoch}: image(s) confirmed STUCK "
|
||||
f"(persistently hard, not improving — check for bad/mislabeled data): "
|
||||
f"{names} — throttling LR x{self.throttle_mult}")
|
||||
# Which resolution is the worst offender for a key, as of the most recent window
|
||||
# it actually has data for — a diagnostic hint, not part of the classification
|
||||
# itself (stuck is still decided on the pooled residual across every resolution).
|
||||
# Uses the LATEST window with data for that key rather than requiring THIS exact
|
||||
# window, since a key can cross the stuck-vote threshold on a window where it
|
||||
# wasn't drawn at all (shuffle variance) — requiring an exact match left the tag
|
||||
# blank most of the time. Only meaningful with >1 resolution in play.
|
||||
by_key_res_by_epoch: dict[str, dict[int, dict[int, list[float]]]] = {}
|
||||
if len(res_sum) > 1:
|
||||
for key, ep, b, res, loss in self._records:
|
||||
by_key_res_by_epoch.setdefault(key, {}).setdefault(ep, {}).setdefault(
|
||||
res, []).append(loss - bmean.get((b, res), 0.0))
|
||||
|
||||
def tag(k: str) -> str:
|
||||
by_epoch = by_key_res_by_epoch.get(k)
|
||||
if not by_epoch:
|
||||
return os.path.basename(k)
|
||||
latest_epoch = max(by_epoch)
|
||||
per_res = by_epoch[latest_epoch]
|
||||
worst_res = max(per_res, key=lambda r: sum(per_res[r]) / len(per_res[r]))
|
||||
return f"{worst_res}px {os.path.basename(k)}"
|
||||
|
||||
names = ", ".join(tag(k) for k in sorted(added))
|
||||
self._log(f"[adaptive-lr] window {epoch}: confirmed stuck (persistently hard, not "
|
||||
f"improving) - check caption if remains stuck: {names} — LR x{self.throttle_mult}")
|
||||
if removed:
|
||||
# No resolution info here — which resolution was worst is only useful while
|
||||
# deciding whether to go inspect the image, not once it's already cleared.
|
||||
names = ", ".join(os.path.basename(k) for k in sorted(removed))
|
||||
self._log(f"[adaptive-lr] window {epoch}: no longer stuck: {names}")
|
||||
self._last_reported_stuck = set(self._confirmed_stuck)
|
||||
|
|
@ -305,8 +361,7 @@ class PerImageAdaptiveLR:
|
|||
"last_reported_stuck": sorted(self._last_reported_stuck),
|
||||
"healthy_epochs": dict(self._healthy_epochs),
|
||||
"retired": sorted(self._retired),
|
||||
"bucket_sum": list(self._bsum),
|
||||
"bucket_cnt": list(self._bcnt),
|
||||
"bucket_stats": [[b, res, self._bsum[(b, res)], self._bcnt[(b, res)]] for (b, res) in self._bsum],
|
||||
# per-key-epoch residual means, flattened as "key\x1fepoch" -> mean, so a resumed run
|
||||
# can keep computing the trend-window improve test across the resume boundary.
|
||||
"per_key_epoch_residual": flattened,
|
||||
|
|
@ -315,10 +370,10 @@ class PerImageAdaptiveLR:
|
|||
def _all_per_key_epoch_residual(self) -> dict:
|
||||
"""Every (key, epoch) mean residual known to this watcher — both from live records this
|
||||
run and any already-restored from a previous checkpoint. Used only when SAVING."""
|
||||
bmean = [self._bsum[i] / self._bcnt[i] if self._bcnt[i] else 0.0 for i in range(_N_BUCKETS)]
|
||||
bmean = {bk: self._bsum[bk] / self._bcnt[bk] for bk in self._bsum if self._bcnt[bk]}
|
||||
out: dict[str, dict[int, list[float]]] = {}
|
||||
for key, ep, b, loss in self._records:
|
||||
out.setdefault(key, {}).setdefault(ep, []).append(loss - bmean[b])
|
||||
for key, ep, b, res, loss in self._records:
|
||||
out.setdefault(key, {}).setdefault(ep, []).append(loss - bmean.get((b, res), 0.0))
|
||||
merged = {k: {ep: sum(v) / len(v) for ep, v in eps.items()} for k, eps in out.items()}
|
||||
for key, eps in self._restored_residuals.items():
|
||||
dest = merged.setdefault(key, {})
|
||||
|
|
@ -341,8 +396,12 @@ class PerImageAdaptiveLR:
|
|||
self._last_reported_stuck = set(state.get("last_reported_stuck", []))
|
||||
self._healthy_epochs = dict(state.get("healthy_epochs", {}))
|
||||
self._retired = set(state.get("retired", []))
|
||||
self._bsum = list(state.get("bucket_sum", self._bsum))
|
||||
self._bcnt = list(state.get("bucket_cnt", self._bcnt))
|
||||
self._bsum = {}
|
||||
self._bcnt = {}
|
||||
for b, res, s, c in state.get("bucket_stats", []):
|
||||
bk = (int(b), int(res))
|
||||
self._bsum[bk] = float(s)
|
||||
self._bcnt[bk] = int(c)
|
||||
# Un-flatten into (key -> epoch -> residual). These are merged into per_key_epoch_mean
|
||||
# at each epoch_boundary() call, never re-derived through a bucket lookup (the bucket
|
||||
# means that produced them belong to the PRIOR run and aren't reconstructable here).
|
||||
|
|
|
|||
Loading…
Reference in New Issue