d=[set() for _ in range(26)]
for idea in ideas:
pre=ord(idea[0])-ord('a')
suf=idea[1:]
d[pre].add(suf)
res=0
for i in range(26):
for j in range(26):
if i!=j:
m=len(d[i].intersection(d[j]))
res+=(len(d[i])-m)*(len(d[j])-m)
return res
어케 풀지 감은 오는데 구현이 막혀서 참고했습니다