Views
No views yet
| Size | CoSQA | AdvTest | CSN-Py | CSN-Ja | CSN-JS | CSN-PHP | CSN-Go | CSN-Ruby | Avg | |
|---|---|---|---|---|---|---|---|---|---|---|
| OpenAI-Embedding-Ada-002 | Unknown | 0.4423 | 0.3808 | 0.6802 | 0.7149 | 0.6750 | 0.6062 | 0.8563 | 0.7472 | 0.6378 |
| OpenAI-Text-embedding-3-large | Unknown | 0.5538 | 0.4684 | 0.7084 | 0.7292 | 0.6813 | 0.5959 | 0.8764 | 0.7525 | 0.6707 |
| jina-embeddings-v2-base-code | 161M | 0.6837 | 0.385 | 0.6634 | 0.6803 | 0.6304 | 0.5701 | 0.8595 | 0.7095 | 0.6477 |
| CodeSage-large | 1.3B | 0.4753 | 0.5267 | 0.7077 | 0.7021 | 0.695 | 0.6133 | 0.8371 | 0.7192 | 0.6595 |
| CodeFuse-CGE-Small | 3.8B | 0.5619 | 0.4639 | 0.6958 | 0.6863 | 0.6564 | 0.6133 | 0.8637 | 0.7341 | 0.6594 |
| OASIS-code-1.5B | 1.5B | 0.5577 | 0.5727 | 0.7369 | 0.7397 | 0.6980 | 0.6384 | 0.8821 | 0.7547 | 0.6975 |
1pip install -U torch
2pip install -U transformers1import torch
2import torch.nn.functional as F
3from torch import Tensor
4from transformers import AutoModel, AutoTokenizer
5def last_token_pool(last_hidden_states: Tensor, attention_mask: Tensor) -> Tensor:
6 left_padding = (attention_mask[:, -1].sum() == attention_mask.shape[0])
7 if left_padding:
8 return last_hidden_states[:, -1]
9 else:
10 sequence_lengths = attention_mask.sum(dim=1) - 1
11 batch_size = last_hidden_states.shape[0]
12 return last_hidden_states[torch.arange(batch_size, device=last_hidden_states.device), sequence_lengths]
13# Add query prompt
14def get_query_prompt(query: str):
15 query_description = 'Given a code search query, retrieve relevant code snippet that answer the query'
16 prompt = f'Instruct: {query_description}\nQuery: {query}'
17 return prompt
18query = "How to do quicksort in python?"
19
20code1 = """def bubble_sort(arr):
21 n = len(arr)
22 for i in range(n):
23 swapped = False
24 for j in range(1, n - i):
25 if arr[j - 1] > arr[j]:
26 arr[j - 1], arr[j] = arr[j], arr[j - 1]
27 swapped = True
28 if not swapped:
29 break
30 return arr"""
31code2 = """def quick_sort(arr):
32 if len(arr) <= 1:
33 return arr
34 else:
35 pivot = arr[0]
36 less = [x for x in arr[1:] if x <= pivot]
37 greater = [x for x in arr[1:] if x > pivot]
38 return quick_sort(less) + [pivot] + quick_sort(greater)"""
39model = AutoModel.from_pretrained("Kwaipilot/OASIS-code-1.5B", output_hidden_states=True)
40tokenizer = AutoTokenizer.from_pretrained("Kwaipilot/OASIS-code-1.5B")
41
42# Tokenize and inference
43inputs = tokenizer([get_query_prompt(query), code1, code2], max_length=1024, padding=True, truncation=True, return_tensors='pt')
44outputs = model(**inputs)
45# Last token pooling
46embeddings = last_token_pool(outputs.hidden_states[-1], inputs['attention_mask'])
47print(embeddings.shape)
48# torch.Size([3, 1536])
49embeddings = F.normalize(embeddings, dim=1, p=2)
50similarity = embeddings @ embeddings.T
51print(similarity[0, 1:])
52# tensor([0.6895, 0.8240])pip install -U sentence-transformers1from sentence_transformers import SentenceTransformer
2# Download from the 🤗 Hub
3model = SentenceTransformer("Kwaipilot/OASIS-code-1.5B")#, model_kwargs={"torch_dtype": torch.bfloat16})
4query = "How to do quicksort in python?"
5code1 = """def bubble_sort(arr):
6 n = len(arr)
7 for i in range(n):
8 swapped = False
9 for j in range(1, n - i):
10 if arr[j - 1] > arr[j]:
11 arr[j - 1], arr[j] = arr[j], arr[j - 1]
12 swapped = True
13 if not swapped:
14 break
15 return arr"""
16code2 = """def quick_sort(arr):
17 if len(arr) <= 1:
18 return arr
19 else:
20 pivot = arr[0]
21 less = [x for x in arr[1:] if x <= pivot]
22 greater = [x for x in arr[1:] if x > pivot]
23 return quick_sort(less) + [pivot] + quick_sort(greater)"""
24# Run inference
25query_embedding = model.encode([query], prompt_name="query")
26code_embeddings = model.encode([code1, code2])
27print(code_embeddings.shape)
28# (2, 1536)
29# Get the similarity scores for the embeddings
30print(model.similarity(query_embedding[0], code_embeddings[0]))
31print(model.similarity(query_embedding[0], code_embeddings[1]))
32# tensor([[0.6895]])
33# tensor([[0.8240]])1@misc{kwaipilotoasis,
2 title = {Optimized Augmentation Strategy for Improved code Search},
3 author = {Kwaipilot team},
4 year = {2024},
5}