(args, model, tokenizer, dev, prune_n=0, prune_m=0)
| 1908 | |
| 1909 | @torch.no_grad() |
| 1910 | def prune_ablate(args, model, tokenizer, dev, prune_n=0, prune_m=0): |
| 1911 | ## SparseGPT code available at: https://github.com/IST-DASLab/sparsegpt/tree/f5c25005a61f96a0933ca2f95705a963585aafaa |
| 1912 | print("Starting ...") |
| 1913 | dataloader, _ = get_loaders( |
| 1914 | "wikitext2", |
| 1915 | nsamples=args.nsamples, |
| 1916 | seed=args.seed, |
| 1917 | seqlen=model.seqlen, |
| 1918 | tokenizer=tokenizer, |
| 1919 | ) |
| 1920 | # dataloader, _ = get_loaders("c4",nsamples=args.nsamples,seed=args.seed,seqlen=model.seqlen,tokenizer=tokenizer) |
| 1921 | |
| 1922 | use_cache = model.config.use_cache |
| 1923 | model.config.use_cache = False |
| 1924 | layers = model.model.layers |
| 1925 | |
| 1926 | if "model.embed_tokens" in model.hf_device_map: |
| 1927 | dev = model.hf_device_map["model.embed_tokens"] |
| 1928 | |
| 1929 | dtype = next(iter(model.parameters())).dtype |
| 1930 | inps = torch.zeros( |
| 1931 | (args.nsamples, model.seqlen, model.config.hidden_size), dtype=dtype, device=dev |
| 1932 | ) |
| 1933 | cache = {"i": 0, "attention_mask": None, "position_ids": None} |
| 1934 | |
| 1935 | class Catcher(nn.Module): |
| 1936 | def __init__(self, module): |
| 1937 | super().__init__() |
| 1938 | self.module = module |
| 1939 | |
| 1940 | def forward(self, inp, **kwargs): |
| 1941 | inps[cache["i"]] = inp |
| 1942 | cache["i"] += 1 |
| 1943 | cache["attention_mask"] = kwargs["attention_mask"] |
| 1944 | cache["position_ids"] = kwargs["position_ids"] |
| 1945 | raise ValueError |
| 1946 | |
| 1947 | layers[0] = Catcher(layers[0]) |
| 1948 | for batch in dataloader: |
| 1949 | try: |
| 1950 | model(batch[0].to(dev)) |
| 1951 | except ValueError: |
| 1952 | pass |
| 1953 | layers[0] = layers[0].module |
| 1954 | torch.cuda.empty_cache() |
| 1955 | |
| 1956 | outs = torch.zeros_like(inps) |
| 1957 | attention_mask = cache["attention_mask"] |
| 1958 | position_ids = cache["position_ids"] |
| 1959 | |
| 1960 | print("Ready.") |
| 1961 | |
| 1962 | for i in range(len(layers)): |
| 1963 | layer = layers[i] |
| 1964 | if f"model.layers.{i}" in model.hf_device_map: |
| 1965 | dev = model.hf_device_map[f"model.layers.{i}"] |
| 1966 | print(f"layer {i} device {dev}") |
| 1967 | inps, outs, attention_mask, position_ids = ( |
no test coverage detected