MCPcopy Create free account
hub / github.com/NVlabs/DiffusionNFT / TextPromptDataset

Class TextPromptDataset

scripts/train_nft_sd3.py:79–95  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

77
78
79class TextPromptDataset(Dataset):
80 def __init__(self, dataset, split="train"):
81 self.file_path = os.path.join(dataset, f"{split}.txt")
82 with open(self.file_path, "r") as f:
83 self.prompts = [line.strip() for line in f.readlines()]
84
85 def __len__(self):
86 return len(self.prompts)
87
88 def __getitem__(self, idx):
89 return {"prompt": self.prompts[idx], "metadata": {}}
90
91 @staticmethod
92 def collate_fn(examples):
93 prompts = [example["prompt"] for example in examples]
94 metadatas = [example["metadata"] for example in examples]
95 return prompts, metadatas
96
97
98class GenevalPromptDataset(Dataset):

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected