CivArchive
    ← All articles
    Published May 20, 2025by AbstractPhila

    Training the T5 Small to be a lora for diffusion models.

    98 views22 reactions2 comments on CivitAI2 collected
    ML Research

    Why?

    It's tiny. 60m parameters and it can be easily fitted into... well anything.

    It's insanely easy to lobotomize though, so I've devised a multitude of methods to teach this model new behaviors without squashing it like a bug.

    WHY THOUGH!?!?!

    Guidance. It's being taught, the rule of 3. It's going to learn, exactly how to guide the image generation in a way that no guidance tool has ever accomplished before.

    https://huggingface.co/AbstractPhil/T5-Small-Human-Attentive-Try2-Pass3
    Not only that, but... it's working. It's adhering internally to large amounts of data; learning exponential amounts of methods to normalize these captions into solidified rule of 3 depictions based on object association, reason, and more.

    The input goes in, the output is starting to diverge from it's echo.

    Essentially, I've devised a way to train it's math, without destroying it.

    The T5 small is a powerful tool.

    You can teach it almost anything in a short period of time, and it will be capable of using that anything with almost no training steps.

    If you train it in the wrong way though, it'll fall apart like a house of cards on the roof of a bullet train.

    Below is a colab drag and drop - required to set the output repo for uploading to a huggingface repo, or removing the lines that upload to repo - blueprint for training the T5-small's internals in a robust and careful methodology.

    I'm essentially DESTROYING it's capability currently. So with this, it's reshaping the internals to conform to the goals of... captions.

    The t5 small is trained with 3 forms of translation and 1 command for "summarize: ", which are all inherently incapable of their baseline jobs.

    However, I've given it a new purpose; caption. This, is currently defining the rule of 3 pathways that will be trainable into diffusion models soon enough.

    This version of the T5-small is different. As it's repurposed, it'll become more and more fluent at these pathways, devising more and more complex methodologies based on it's internals - forming even more accurate guidance than before in more complex methodologies and ways.

    Given the drop of adherence and the reduction of loss reaching a certain point, with the multi-prompt response from BLEU reaching a certain peak point; this model will be ready to trained into SDXL as a flat guidance spine lora.

    What does it mean though?

    I'm teaching the T5 to be a solid guidance structure for any diffusion model; as a supplement to the existing encoders. Like a cybernetic limb, or an ear, or in this case; a part of the brain implanted in a way that guides the vectors in a carefully utilizable way.

    It means... it'll guide SDXL in a way that has never been guided before; both speeding up the inference, improving fidelity, quality, and more; all because the baseline SDXL variations will be allowed to focus more of their internal mechanisms ON THEIR OWN GOALS.

    It's like having an extra set of eyes helping guide the output - not through a blanket mathematics response, but through an interpolated learned response that can simply be snapped on like a lora.

    # train_t5_small_human_attentive_full_pipeline.py
    
    import os, random, csv, collections
    from pathlib import Path
    
    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    from torch.utils.data import IterableDataset
    
    from transformers import (
        T5ForConditionalGeneration,
        T5TokenizerFast,
        DataCollatorForSeq2Seq,
        TrainingArguments,
        Trainer,
        TrainerCallback,
        get_cosine_schedule_with_warmup,
    )
    from huggingface_hub import login, hf_hub_download
    from sacrebleu import corpus_bleu
    
    # ──────────────────────────────────────────────────────────
    # 🔐 Login to HF Hub
    # ──────────────────────────────────────────────────────────
    from google.colab import userdata
    
    os.environ["HF_TOKEN"] = userdata.get('HF_TOKEN')
    login(token=os.getenv("HF_TOKEN"))
    
    # ──────────────────────────────────────────────────────────
    # ⚙️ Model & Tokenizer (+ higher dropout)
    # ──────────────────────────────────────────────────────────
    model = T5ForConditionalGeneration.from_pretrained("AbstractPhil/T5-Small-Human-Attentive-Try2-Pass2")
    tokenizer = T5TokenizerFast.from_pretrained("t5-small")
    
    # bump dropout for heavy‐LR regime (Tier 6)
    model.config.dropout_rate = 0.3
    model.config.attention_dropout = 0.3
    
    # Task prefix config
    model.config.task_specific_params = {
        "caption": {
            "prefix":         "caption: ",
            "max_length":     64,
            "num_beams":      4,
            "length_penalty": 1.0,
            "early_stopping": True,
        }
    }
    
    # unfreeze everything
    for p in model.parameters(): p.requires_grad = True
    
    # ──────────────────────────────────────────────────────────
    # ░ Tier 2 ░ WeightScheduler “slow-bleed”
    # ──────────────────────────────────────────────────────────
    class WeightScheduler:
        def __init__(self, total_steps, high_start=5.0, low_start=0.2, end=1.0):
            """
            total_steps: number of steps over which to decay extremes → 1.0
            high_start: initial multiplier for 'high' tokens
            low_start:  initial multiplier for 'low' tokens
            end:        final multiplier for both
            """
            self.total_steps = total_steps
            self.high_start, self.low_start, self.end = high_start, low_start, end
    
        def get_weight(self, step, token_class):
            """
            step: current global step
            token_class: 0=neutral, 1=high, 2=low (scalar or tensor)
            returns: same shape as token_class, float weights
            """
            frac = min(step / self.total_steps, 1.0)
            high_w = self.high_start + frac * (self.end - self.high_start)
            low_w  = self.low_start  + frac * (self.end - self.low_start)
            # branch on class
            if isinstance(token_class, torch.Tensor):
                w = torch.ones_like(token_class, dtype=torch.float, device=token_class.device)
                w = torch.where(token_class == 1, high_w, w)
                w = torch.where(token_class == 2, low_w,  w)
                return w
            else:
                return {0:1.0, 1:high_w, 2:low_w}[int(token_class)]
    
    # ──────────────────────────────────────────────────────────
    # ░ Tier 4 ░ MaskRescaler
    # ──────────────────────────────────────────────────────────
    class MaskRescaler:
        def __init__(self, window=2, factor=1.2):
            """
            window: number of tokens on each side of a masked position
            factor: multiplier for neighbor weights
            """
            self.window, self.factor = window, factor
    
        def rescale(self, weights, mask_positions):
            """
            weights: [B, L] tensor of per-token weights
            mask_positions: list of lists of masked indices per example
            """
            B, L = weights.shape
            for b, positions in enumerate(mask_positions):
                for idx in positions:
                    left  = max(0, idx - self.window)
                    right = min(L, idx + self.window + 1)
                    weights[b, left:right] *= self.factor
            return weights
    
    # ──────────────────────────────────────────────────────────
    # ░ Tier 5 ░ MixedCaptionDataset (mask+half-swap)
    # ──────────────────────────────────────────────────────────
    class MixedCaptionDataset(IterableDataset):
        def __init__(self,
                     tokenizer,
                     repo_id: str,
                     examples_per_file: int = 5_000_000,
                     num_files: int = 10,
                     batch_size: int = 256,
                     low_q: float = 0.4,      # mid-quantile mask
                     high_q: float = 0.6,
                     decay_step: int = 1000,  # decay masked counts
                     swap_prob: float = 0.5   # Tier 5 mixing
        ):
            self.tokenizer   = tokenizer
            self.repo_id     = repo_id
            self.files       = [f"captions/caption_{i+2:03d}.csv" for i in range(num_files)]
            self.batch_size  = batch_size
            self.low_q, self.high_q, self.decay_step = low_q, high_q, decay_step
            self.swap_prob   = swap_prob
            self.examples    = examples_per_file * num_files
    
            # running global counts of label‐tokens
            self.token_counts = collections.Counter()
    
        def __len__(self):
            return self.examples
    
        def __iter__(self):
            # stream each file, batch‐tokenize, update counts & mask
            for rel in self.files:
                path = hf_hub_download(self.repo_id, rel, repo_type="dataset")
                with open(path, encoding="utf-8") as f:
                    reader, buf = csv.DictReader(f), []
                    for row in reader:
                        text = row.get("text","").strip()
                        if not text: continue
                        buf.append(("caption: "+text, text))
                        if len(buf) >= self.batch_size:
                            for sample in self._batch_tokenize(buf):
                                yield sample
                            buf.clear()
                    if buf:
                        for sample in self._batch_tokenize(buf):
                            yield sample
    
        def _batch_tokenize(self, batch):
            inputs, targets = zip(*batch)
            enc = self.tokenizer(
                list(inputs),
                padding="max_length", truncation=True, max_length=256,
                return_tensors="pt"
            )
            lbl = self.tokenizer(
                list(targets),
                padding="max_length", truncation=True, max_length=64,
                return_tensors="pt"
            )
    
            # 1) update global counts
            flat = lbl.input_ids.view(-1).tolist()
            valid = [tok for tok in flat if tok != self.tokenizer.pad_token_id]
            self.token_counts.update(valid)
    
            # 2) select mid-quantile to mask
            items = sorted(self.token_counts.items(), key=lambda x: x[1])
            n = len(items)
            if n:
                toks, cnts = zip(*items)
                low_i  = int(n * self.low_q)
                high_i = int(n * self.high_q)
                mask_set = set(toks[low_i:high_i])
            else:
                mask_set = set()
    
            # 3) decay their counts
            for t in mask_set:
                self.token_counts[t] = max(0, self.token_counts[t] - self.decay_step)
    
            # 4) now build samples, either masked or half-swap
            samples = []
            pad = self.tokenizer.pad_token_id
            for i in range(len(batch)):
                label_ids = lbl.input_ids[i]
                if random.random() < self.swap_prob:
                    # we'll handle swapping higher up: mark with special flag
                    samples.append({
                        "input_ids":      enc.input_ids[i].tolist(),
                        "attention_mask": enc.attention_mask[i].tolist(),
                        "labels":         lbl.input_ids[i].tolist(),
                        "_swap":          True
                    })
                else:
                    # mask mid-quantile tokens
                    mask_bool = torch.tensor([tok in mask_set for tok in label_ids])
                    masked    = torch.where(mask_bool, -100, label_ids)
                    samples.append({
                        "input_ids":      enc.input_ids[i].tolist(),
                        "attention_mask": enc.attention_mask[i].tolist(),
                        "labels":         masked.tolist(),
                        "_swap":          False
                    })
            # if any swaps, pair them sequentially
            out = []
            i = 0
            while i < len(samples):
                s = samples[i]
                if s["_swap"] and i+1 < len(samples):
                    t = samples[i+1]
                    # half-swap inputs & labels between s and t
                    i_len = len(s["input_ids"])
                    l_len = len(s["labels"])
                    hi = i_len//2; hl = l_len//2
                    i1 = torch.tensor(s["input_ids"]); i2 = torch.tensor(t["input_ids"])
                    l1 = torch.tensor(s["labels"]);     l2 = torch.tensor(t["labels"])
                    # swap halves
                    new_i1 = torch.cat([i1[:hi], i2[hi:]], dim=0)
                    new_i2 = torch.cat([i2[:hi], i1[hi:]], dim=0)
                    new_l1 = torch.cat([l1[:hl], l2[hl:]], dim=0)
                    new_l2 = torch.cat([l2[:hl], l1[hl:]], dim=0)
                    for rec, ni, nl in [(s,new_i1,new_l1),(t,new_i2,new_l2)]:
                        rec["input_ids"]      = ni.tolist()
                        rec["attention_mask"] = ((ni!=pad).long()).tolist()
                        rec["labels"]         = nl.tolist()
                    out.extend([s,t])
                    i += 2
                else:
                    out.append(s)
                    i += 1
            return out
    
    # ──────────────────────────────────────────────────────────
    # ░ Tier 2/3/4 ░ WeightedTrainer
    # ──────────────────────────────────────────────────────────
    class WeightedTrainer(Trainer):
        def __init__(self, *args,
                     weight_scheduler: WeightScheduler,
                     mask_rescaler:    MaskRescaler,
                     low_q_extreme:  float = 0.1,
                     high_q_extreme: float = 0.9,
                     **kwargs):
            super().__init__(*args, **kwargs)
            # inject our scheduler & rescaler
            self.weight_scheduler = weight_scheduler
            self.mask_rescaler    = mask_rescaler
            self.low_q_extreme    = low_q_extreme
            self.high_q_extreme   = high_q_extreme
    
            # precompute total steps for scheduler
            if self.args.max_steps > 0:
                total = self.args.max_steps
            else:
                # approximate
                ds_len = len(self.train_dataset)
                bs     = self.args.per_device_train_batch_size * self.args.gradient_accumulation_steps
                total  = int(ds_len/bs) * self.args.num_train_epochs
            self.weight_scheduler.total_steps = total
    
        def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
            # 1) Remove labels and our custom swap flag
            labels = inputs.pop("labels")
            inputs.pop("_swap", None)
    
            # 2) Forward pass
            outputs = model(**inputs, labels=labels)
            logits  = outputs.logits       # shape: [B, L, V]
            B, L, V = logits.size()
    
            # 3) Flatten labels and select valid positions (ignore_index = -100)
            flat_labels = labels.view(-1)
            valid_mask  = flat_labels != -100
            valid_idx   = valid_mask.nonzero(as_tuple=True)[0]
            valid_lbls  = flat_labels[valid_idx]
    
            # 4) Get current global step
            step = self.state.global_step
    
            # ── Tier 2: classify tokens into high/low/neutral based on global counts ──
            gc    = self.train_dataset.token_counts
            items = sorted(gc.items(), key=lambda x: x[1])
            n     = len(items)
            low_i  = int(n * self.low_q_extreme)
            high_i = int(n * self.high_q_extreme)
            low_set  = set(tok for tok, _ in items[:low_i])
            high_set = set(tok for tok, _ in items[high_i:])
            # Build a class vector: 0=neutral, 1=high, 2=low
            classes = torch.zeros_like(flat_labels)
            cls_list = []
            for tid in valid_lbls.tolist():
                if tid in high_set:
                    cls_list.append(1)
                elif tid in low_set:
                    cls_list.append(2)
                else:
                    cls_list.append(0)
            classes[valid_idx] = torch.tensor(cls_list, device=classes.device)
    
            # 5) Tier 2: compute slow-bleed weights
            bleed_w = self.weight_scheduler.get_weight(step, classes)
    
            # ── Tier 3: inverse-frequency weighting (per-batch) ──
            uniq, cnts = valid_lbls.unique(return_counts=True)
            inv        = cnts.max().float() / cnts.float()
            inv_freq   = torch.ones(V, device=logits.device)
            inv_freq[uniq] = inv
            inv_w_flat = torch.ones_like(flat_labels, dtype=logits.dtype)
            inv_w_flat[valid_idx] = inv_freq[valid_lbls]
    
            # 6) Combine bleed + inverse-frequency
            w_flat  = bleed_w * inv_w_flat
            weights = w_flat.view(B, L)
    
            # ── Tier 4: mask-triggered local rescaling ──
            mask_pos = [(labels[b] == -100).nonzero(as_tuple=True)[0].tolist() for b in range(B)]
            weights  = self.mask_rescaler.rescale(weights, mask_pos)
    
            # 7) Compute weighted negative log-likelihood loss
            logprobs = F.log_softmax(logits, dim=-1).view(-1, V)
            sel_lp   = logprobs[valid_idx, valid_lbls]
            w_sel    = w_flat[valid_idx]
            loss     = -(w_sel * sel_lp).sum() / w_sel.sum()
    
            return (loss, outputs) if return_outputs else loss
    
    
    # ──────────────────────────────────────────────────────────
    # ░ Tier 7 ░ BLEU Eval Callback
    # ──────────────────────────────────────────────────────────
    class BLEUCallback(TrainerCallback):
        def __init__(self, model, tokenizer, references, prompts, every_n_steps=500):
            self.model     = model
            self.tokenizer = tokenizer
            self.refs      = references  # list of list of str
            self.prompts   = prompts     # list of str
            self.every_n   = every_n_steps
            self.metric    = corpus_bleu
    
        def on_step_end(self, args, state, control, **kwargs):
            if state.global_step and state.global_step % self.every_n == 0:
                self.model.eval()
                preds = []
                for p in self.prompts:
                    batch = self.tokenizer(
                        "caption: " + p,
                        truncation=True, padding="max_length",
                        max_length=256, return_tensors="pt"
                    ).to(self.model.device)
                    out = self.model.generate(
                        **batch,
                        max_length=64,
                        num_beams=4,
                        early_stopping=True,
                    )
                    dec = self.tokenizer.decode(out[0], skip_special_tokens=True)
                    preds.append(dec)
                score = self.metric(preds, [self.refs]).score
                print(f"▶️  Step {state.global_step}: mixed‐caption BLEU = {score:.2f}")
                self.model.train()
    
    # ──────────────────────────────────────────────────────────
    # 🧠 Assemble & Launch
    # ──────────────────────────────────────────────────────────
    
    # held‐out mix references & prompts for BLEU
    # held‐out mix references & prompts for BLEU
    mix_prompts = [
        "a room of tacos",
        "a forest of neon mushrooms",
        "a beach covered in colorful seashells",
        "a mountain peak above the clouds",
        "a desk cluttered with vintage cameras",
        "a city skyline at dusk",
        "a cat sleeping on a windowsill",
        "a bouquet of wildflowers in a mason jar",
        "a red sports car on a desert road",
        "a steaming cup of coffee on a book",
        "a group of hot air balloons",
        "a snowy village at night",
        "a winding forest path in autumn",
        "a pair of sneakers by the door",
        "a stack of pancakes with syrup",
        "a woodland stream with stones",
        "a baby elephant playing in mud",
        "a painted guitar leaning on a chair",
        "a row of colorful umbrellas on a beach",
        "a vintage typewriter on a desk",
    ]
    
    mix_references = [
        "brightly lit room filled entirely with tacos",
        "an enchanted forest glowing under neon mushrooms",
        "a sunlit beach strewn with a rainbow of seashells",
        "a solitary mountain peak rising above a sea of clouds",
        "a wooden desk overflowing with vintage film cameras",
        "a city skyline silhouetted against a dusky purple sky",
        "a fluffy cat curled up on a sunlit windowsill",
        "a mason jar brimming with wildflowers picked from a meadow",
        "a sleek red sports car speeding along an empty desert highway",
        "a steaming cup of coffee resting on an open leather-bound book",
        "a colorful fleet of hot air balloons drifting across a clear sky",
        "a quiet snowy village illuminated by warm street lamps",
        "a winding forest trail blanketed in golden autumn leaves",
        "a worn pair of sneakers casually placed by the wooden door",
        "a stack of fluffy pancakes drenched in maple syrup",
        "a gentle woodland stream flowing over moss-covered stones",
        "a baby elephant joyfully splashing mud with its trunk",
        "a brightly painted guitar leaning against a wooden chair",
        "a line of vibrant umbrellas casting shade on the sandy beach",
        "a retro typewriter sitting on a cluttered writing desk",
    ]
    
    
    dataset  = MixedCaptionDataset(
        tokenizer,
        repo_id="AbstractPhil/human-templated-captions-1b",
        examples_per_file=5_000_000,
        num_files=2,
        batch_size=256,
        low_q=0.4, high_q=0.6,
        decay_step=1000,
        swap_prob=0.5,
    )
    collator = DataCollatorForSeq2Seq(tokenizer, model=model, label_pad_token_id=-100)
    
    # build scheduler + rescaler
    # we’ll fill total_steps after Trainer init
    dummy_total = 1_000_000
    weight_scheduler = WeightScheduler(total_steps=dummy_total,
                                       high_start=5.0, low_start=0.2, end=1.0)
    mask_rescaler    = MaskRescaler(window=2, factor=1.2)
    
    training_args = TrainingArguments(
        output_dir                 = "./checkpoints-full",
        per_device_train_batch_size=256,
        gradient_accumulation_steps=4,
        dataloader_num_workers     =6,
        learning_rate              =1e-3,
        optim                      ="adafactor",
        weight_decay               =0.01,
        lr_scheduler_type          ="cosine",
        warmup_steps               =500,
        num_train_epochs           =4,
        logging_steps              =50,
        save_steps                 =500,
        save_total_limit           =5,
        max_grad_norm              =1.0,           
        bf16                       =True,
        remove_unused_columns      =False,
        push_to_hub                =True,
        hub_model_id               ="AbstractPhil/T5-Small-Human-Attentive-Try2-Pass3",
        report_to                  ="none",
    )
    
    trainer = WeightedTrainer(
        model            =model,
        args             =training_args,
        train_dataset    =dataset,
        data_collator    =collator,
        callbacks        =[
            BLEUCallback(model, tokenizer, mix_references, mix_prompts, every_n_steps=50)
        ],
        weight_scheduler =weight_scheduler,
        mask_rescaler    =mask_rescaler,
        low_q_extreme    =0.1,
        high_q_extreme   =0.9,
    )
    
    if __name__ == "__main__":
        trainer.train()
        trainer.push_to_hub()