(args, model, tokenizer, dev, prune_n=0, prune_m=0)
| 1789 | |
| 1790 | |
| 1791 | def prune_sparsegpt(args, model, tokenizer, dev, prune_n=0, prune_m=0): |
| 1792 | ## SparseGPT code available at: https://github.com/IST-DASLab/sparsegpt/tree/f5c25005a61f96a0933ca2f95705a963585aafaa |
| 1793 | print("Starting ...") |
| 1794 | dataloader, _ = get_loaders( |
| 1795 | "wikitext2", |
| 1796 | nsamples=args.nsamples, |
| 1797 | seed=args.seed, |
| 1798 | seqlen=model.seqlen, |
| 1799 | tokenizer=tokenizer, |
| 1800 | ) |
| 1801 | # dataloader, _ = get_loaders("c4",nsamples=args.nsamples,seed=args.seed,seqlen=model.seqlen,tokenizer=tokenizer) |
| 1802 | |
| 1803 | use_cache = model.config.use_cache |
| 1804 | model.config.use_cache = False |
| 1805 | layers = model.model.layers |
| 1806 | |
| 1807 | if "model.embed_tokens" in model.hf_device_map: |
| 1808 | dev = model.hf_device_map["model.embed_tokens"] |
| 1809 | |
| 1810 | dtype = next(iter(model.parameters())).dtype |
| 1811 | inps = torch.zeros( |
| 1812 | (args.nsamples, model.seqlen, model.config.hidden_size), dtype=dtype, device=dev |
| 1813 | ) |
| 1814 | cache = {"i": 0, "attention_mask": None, "position_ids": None} |
| 1815 | |
| 1816 | class Catcher(nn.Module): |
| 1817 | def __init__(self, module): |
| 1818 | super().__init__() |
| 1819 | self.module = module |
| 1820 | |
| 1821 | def forward(self, inp, **kwargs): |
| 1822 | inps[cache["i"]] = inp |
| 1823 | cache["i"] += 1 |
| 1824 | cache["attention_mask"] = kwargs["attention_mask"] |
| 1825 | cache["position_ids"] = kwargs["position_ids"] |
| 1826 | raise ValueError |
| 1827 | |
| 1828 | layers[0] = Catcher(layers[0]) |
| 1829 | for batch in dataloader: |
| 1830 | try: |
| 1831 | model(batch[0].to(dev)) |
| 1832 | except ValueError: |
| 1833 | pass |
| 1834 | layers[0] = layers[0].module |
| 1835 | torch.cuda.empty_cache() |
| 1836 | |
| 1837 | outs = torch.zeros_like(inps) |
| 1838 | attention_mask = cache["attention_mask"] |
| 1839 | position_ids = cache["position_ids"] |
| 1840 | |
| 1841 | print("Ready.") |
| 1842 | |
| 1843 | for i in range(len(layers)): |
| 1844 | layer = layers[i] |
| 1845 | if f"model.layers.{i}" in model.hf_device_map: |
| 1846 | dev = model.hf_device_map[f"model.layers.{i}"] |
| 1847 | print(f"layer {i} device {dev}") |
| 1848 | inps, outs, attention_mask, position_ids = ( |
no test coverage detected