Instructions to use xfcghj/AR with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use xfcghj/AR with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("xfcghj/AR", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
| import os | |
| import glob | |
| import tarfile | |
| import pickle | |
| import zstandard | |
| import torch | |
| import torch.distributed as dist | |
| from torch.utils.data import IterableDataset, DataLoader | |
| class DyMeshIterableDataset(IterableDataset): | |
| def __init__(self, data_dir, rank=0, world_size=1): | |
| super(DyMeshIterableDataset, self).__init__() | |
| self.tar_paths = sorted(glob.glob(os.path.join(data_dir, "shard_*.tar.zst"))) | |
| self.my_tar_paths = [p for i, p in enumerate(self.tar_paths) if i % world_size == rank] | |
| print(f"[Rank {rank}] 负责处理 {len(self.my_tar_paths)} 个分片") | |
| def __iter__(self): | |
| worker_info = torch.utils.data.get_worker_info() | |
| if worker_info is None: | |
| curr_paths = self.my_tar_paths | |
| else: | |
| per_worker = (len(self.my_tar_paths) + worker_info.num_workers - 1) // worker_info.num_workers | |
| wid = worker_info.id | |
| curr_paths = self.my_tar_paths[wid * per_worker : (wid + 1) * per_worker] | |
| dctx = zstandard.ZstdDecompressor() | |
| for tar_path in curr_paths: | |
| try: | |
| with open(tar_path, 'rb') as fh: | |
| with dctx.stream_reader(fh) as reader: | |
| with tarfile.open(fileobj=reader, mode='r|') as tar: | |
| for member in tar: | |
| if member.isfile(): | |
| f = tar.extractfile(member) | |
| if f: | |
| data = pickle.load(f) | |
| yield {'vertices': data['vertices'], 'faces': data['faces'], 'caption': data.get('caption', '无')} | |
| except Exception: continue | |
| def get_mesh_dataloader(data_dir, batch_size=4, num_workers=2): | |
| rank = int(os.environ.get("RANK", 0)) | |
| world_size = int(os.environ.get("WORLD_SIZE", 1)) | |
| dataset = DyMeshIterableDataset(data_dir, rank=rank, world_size=world_size) | |
| return DataLoader(dataset, batch_size=batch_size, collate_fn=lambda x: x, num_workers=num_workers, pin_memory=True) | |
| def main(): | |
| dist.init_process_group(backend="nccl") | |
| local_rank = int(os.environ["LOCAL_RANK"]) | |
| global_rank = dist.get_rank() | |
| world_size = dist.get_world_size() | |
| torch.cuda.set_device(local_rank) | |
| dataloader = get_mesh_dataloader("/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/data/shards/", batch_size=40, num_workers=8) | |
| local_batch_count = 0 | |
| local_sample_count = 0 | |
| for batch_idx, batch in enumerate(dataloader): | |
| local_batch_count += 1 | |
| local_sample_count += len(batch) # 统计实际样本数 | |
| if batch_idx % 50 == 0: | |
| print(f"[Rank {global_rank}] 进度: 处理至第 {local_batch_count} 个 batch (累计样本: {local_sample_count})") | |
| dist.barrier() | |
| # 汇总 Batch 数量 | |
| batch_counts = torch.zeros(world_size, dtype=torch.long, device=f"cuda:{local_rank}") | |
| batch_counts[global_rank] = local_batch_count | |
| dist.all_reduce(batch_counts, op=dist.ReduceOp.SUM) | |
| # 汇总 Sample 数量 | |
| sample_counts = torch.zeros(world_size, dtype=torch.long, device=f"cuda:{local_rank}") | |
| sample_counts[global_rank] = local_sample_count | |
| dist.all_reduce(sample_counts, op=dist.ReduceOp.SUM) | |
| if global_rank == 0: | |
| print("\n" + "="*50) | |
| print("📊 分布式训练统计报告") | |
| for i in range(world_size): | |
| print(f" Rank {i}: 处理了 {batch_counts[i].item()} 个 Batch, 共 {sample_counts[i].item()} 个样本") | |
| print("-" * 50) | |
| print(f" 全局总 Batch 数: {batch_counts.sum().item()}") | |
| print(f" 全局总样本数 : {sample_counts.sum().item()}") | |
| print("="*50 + "\n") | |
| dist.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() |